Ë
    þÍ:jº  ã                   ó°   — 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m	Z	 d dl
mZ d dlmZmZ d dlmZ d dlmZ d d	lmZ d d
lmZmZ esdgZ G d„ de«      Zy)é    )ÚSequence)ÚAnyÚListÚOptionalÚUnionN)ÚTensorÚtensor)ÚLiteral)Ú_total_variation_computeÚ_total_variation_update)ÚMetric)Údim_zero_cat)Ú_MATPLOTLIB_AVAILABLE)Ú_AX_TYPEÚ_PLOT_OUT_TYPEzTotalVariation.plotc                   óì   ‡ — e Zd ZU dZdZeed<   dZeed<   dZeed<   dZ	e
ed<   eed	<   ee   ed
<   eed<   ddeed      deddfˆ fd„Zdeddfd„Zdefd„Z	 ddeeeee   f      dee   defd„Zˆ xZS )ÚTotalVariationa¦  Compute Total Variation loss (`TV`_).

    As input to ``forward`` and ``update`` the metric accepts the following input

    - ``img`` (:class:`~torch.Tensor`): A tensor of shape ``(N, C, H, W)`` consisting of images

    As output of `forward` and `compute` the metric returns the following output

    - ``sdi`` (:class:`~torch.Tensor`): if ``reduction!='none'`` returns float scalar tensor with average TV value
      over sample else returns tensor of shape ``(N,)`` with TV values per sample

    Args:
        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

        kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info.

    Raises:
        ValueError:
            If ``reduction`` is not one of ``'sum'``, ``'mean'``, ``'none'`` or ``None``

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

    FÚfull_state_updateTÚis_differentiableÚhigher_is_betterg        Úplot_lower_boundÚnum_elementsÚ
score_listÚscoreÚ	reduction)ÚmeanÚsumÚnoneÚkwargsÚreturnNc                 ó0  •— t        ‰| �  di |¤Ž |�|dvrt        d«      ‚|| _        | j	                  dg d¬«       | j	                  dt        dt        j                  ¬«      d	¬«       | j	                  d
t        dt        j                  ¬«      d	¬«       y )N)r   r   r   zHExpected argument `reduction` to either be 'sum', 'mean', 'none' or Noner   Úcat)ÚdefaultÚdist_reduce_fxr   r   )Údtyper   r   © )	ÚsuperÚ__init__Ú
ValueErrorr   Ú	add_stater	   ÚtorchÚfloatÚint)Úselfr   r   Ú	__class__s      €új/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchmetrics/image/tv.pyr(   zTotalVariation.__init__K   s‚   ø€ Ü‰ÑÑ"˜6Ò"ØÐ  YÐ6MÑ%MÜÐgÓhÐhØ"ˆŒà�‰�|¨RÀˆÔFØ�‰�w¬¨q¼¿¹Ô(DÐUZˆÔ[Ø�‰�~¬v°a¼u¿y¹yÔ/IÐZ_ˆÕ`ó    Úimgc                 óþ   — t        |«      \  }}| j                  �| j                  dk(  r| j                  j                  |«       n#| xj                  |j                  «       z  c_        | xj                  |z  c_        y)z0Update current score with batch of input images.Nr   )r   r   r   Úappendr   r   r   )r.   r2   r   r   s       r0   ÚupdatezTotalVariation.updateU   s]   € ä5°cÓ:Ñˆˆ|Ø�>‰>Ð! T§^¡^°vÒ%=Ø�O‰O×"Ñ" 5Õ)à�JŠJ˜%Ÿ)™)›+Ñ%�JØ×Ò˜\Ñ)Ör1   c                 ó¼   — | j                   �| j                   dk(  rt        | j                  «      n| j                  }t	        || j
                  | j                   «      S )zCompute final total variation.r   )r   r   r   r   r   r   )r.   r   s     r0   ÚcomputezTotalVariation.compute^   sG   € à15·±Ð1GÈ4Ï>É>Ð]cÒKc”˜TŸ_™_Ô-Ðim×isÑisˆÜ'¨¨t×/@Ñ/@À$Ç.Á.ÓQÐQr1   ÚvalÚaxc                 ó&   — | j                  ||«      S )aî  Plot a single or multiple values from the metric.

        Args:
            val: Either a single result from calling `metric.forward` or `metric.compute` or a list of these results.
                If no value is provided, will automatically call `metric.compute` and plot that result.
            ax: An matplotlib axis object. If provided will add plot to that axis

        Returns:
            Figure and Axes object

        Raises:
            ModuleNotFoundError:
                If `matplotlib` is not installed

        .. plot::
            :scale: 75

            >>> # Example plotting a single value
            >>> import torch
            >>> from torchmetrics.image import TotalVariation
            >>> metric = TotalVariation()
            >>> metric.update(torch.rand(5, 3, 28, 28))
            >>> fig_, ax_ = metric.plot()

        .. plot::
            :scale: 75

            >>> # Example plotting multiple values
            >>> import torch
            >>> from torchmetrics.image import TotalVariation
            >>> metric = TotalVariation()
            >>> values = [ ]
            >>> for _ in range(10):
            ...     values.append(metric(torch.rand(5, 3, 28, 28)))
            >>> fig_, ax_ = metric.plot(values)

        )Ú_plot)r.   r8   r9   s      r0   ÚplotzTotalVariation.plotc   s   € ðP �z‰z˜#˜rÓ"Ð"r1   )r   )NN)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   ÚboolÚ__annotations__r   r   r   r,   r   r   r   r
   r   r(   r5   r7   r   r   r   r   r<   Ú__classcell__)r/   s   @r0   r   r      sÖ   ø… ñ ðD $Ð�tÓ#Ø"Ð�tÓ"Ø"Ð�dÓ"Ø!Ð�eÓ!àÓØ�V‘ÓØƒMña (¨7Ð3HÑ+IÑ"Jð aÐ^að aÐfjõ að*˜&ð * Tó *ðR˜ó Rð _cñ(#Ø˜E &¨(°6Ñ*:Ð":Ñ;Ñ<ð(#ØIQÐRZÑI[ð(#à	÷(#r1   r   )Úcollections.abcr   Útypingr   r   r   r   r+   r   r	   Útyping_extensionsr
   Ú torchmetrics.functional.image.tvr   r   Útorchmetrics.metricr   Útorchmetrics.utilities.datar   Útorchmetrics.utilities.importsr   Útorchmetrics.utilities.plotr   r   Ú__doctest_skip__r   r&   r1   r0   ú<module>rM      sB   ðõ %ß -Ó -ã ß  Ý %ç ^Ý &Ý 4Ý @ß @áØ-Ð.Ðôl#�Võ l#r1   