Ë
    þÍ:j@!  ã                   ó¸   — d dl mZ d dlmZ d dl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)Úpartial)ÚAnyÚOptionalÚUnionN)ÚTensorÚtensor)ÚLiteral)Ú_psnr_computeÚ_psnr_update)ÚMetric)Úrank_zero_warn)Ú_MATPLOTLIB_AVAILABLE)Ú_AX_TYPEÚ_PLOT_OUT_TYPEzPeakSignalNoiseRatio.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	<   	 	 	 dd	ee
ee
e
f   f   de
ded   deeeeedf   f      dedd
fˆ fd„Zde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 )ÚPeakSignalNoiseRatioa4  `Compute Peak Signal-to-Noise Ratio`_ (PSNR).

    .. math:: \text{PSNR}(I, J) = 10 * \log_{10} \left(\frac{\max(I)^2}{\text{MSE}(I, J)}\right)

    Where :math:`\text{MSE}` denotes the `mean-squared-error`_ function.

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

    - ``preds`` (:class:`~torch.Tensor`): Predictions from model of shape ``(N,C,H,W)``
    - ``target`` (:class:`~torch.Tensor`): Ground truth values of shape ``(N,C,H,W)``

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

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

    Args:
        data_range:
            the range of the data. If a tuple is provided, then the range is calculated as the difference and
            input is clamped between the values.
        base: a base of a logarithm to use.
        reduction: a method to reduce metric score over labels.

            - ``'elementwise_mean'``: takes the mean (default)
            - ``'sum'``: takes the sum
            - ``'none'`` or ``None``: no reduction will be applied

        dim:
            Dimensions to reduce PSNR scores over, provided as either an integer or a list of integers. Default is
            None meaning scores will be reduced across all dimensions and all batches.
        kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info.

    Example:
        >>> from torchmetrics.image import PeakSignalNoiseRatio
        >>> psnr = PeakSignalNoiseRatio(data_range=3.0)
        >>> preds = torch.tensor([[0.0, 1.0], [2.0, 3.0]])
        >>> target = torch.tensor([[3.0, 2.0], [1.0, 0.0]])
        >>> psnr(preds, target)
        tensor(2.5527)

    TÚis_differentiableÚhigher_is_betterFÚfull_state_updateç        Úplot_lower_boundÚ
data_rangeNÚbaseÚ	reduction)Úelementwise_meanÚsumÚnoneNÚdim.ÚkwargsÚreturnc                 óœ  •— t        ‰| �  di |¤Ž |€|dk7  rt        d|› d�«       |€;| j                  dt	        d«      d¬«       | j                  dt	        d	«      d¬«       n(| j                  dg d
¬«       | j                  dg d
¬«       d | _        t        |t        «      rN| j                  dt	        |d   |d	   z
  «      d¬«       t        t        j                  |d	   |d   ¬«      | _        n&| j                  dt	        t        |«      «      d¬«       || _        || _        t        |t        «      rt        |«      | _        y || _        y )Nr   zThe `reduction=z.` will not have any effect when `dim` is None.Úsum_squared_errorr   r   )ÚdefaultÚdist_reduce_fxÚtotalr   Úcatr   é   Úmean)ÚminÚmax© )ÚsuperÚ__init__r   Ú	add_stater	   Úclamping_fnÚ
isinstanceÚtupler   ÚtorchÚclampÚfloatr   r   r   r   )Úselfr   r   r   r   r    Ú	__class__s         €úl/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchmetrics/image/psnr.pyr.   zPeakSignalNoiseRatio.__init__Q   s'  ø€ ô 	‰ÑÑ"˜6Ò"àˆ;˜9Ð(:Ò:Ü˜_¨Y¨KÐ7eÐfÔgàˆ;Ø�N‰NÐ.¼¸s»ÐTYˆNÔZØ�N‰N˜7¬F°1«IÀeˆNÕLà�N‰NÐ.¸È5ˆNÔQØ�N‰N˜7¨B¸uˆNÔEàˆÔÜ�j¤%Ô(Ø�N‰N˜<´¸
À1¹È
ÐSTÉÑ8UÓ1VÐgmˆNÔnÜ&¤u§{¡{¸
À1¹È:ÐVWÉ=ÔYˆDÕà�N‰N˜<´¼¸jÓ8IÓ1JÐ[aˆNÔbàˆŒ	Ø"ˆŒÜ!+¨C´Ô!:”5˜“:ˆ�Àˆ�ó    ÚpredsÚtargetc                 óH  — | j                   �"| j                  |«      }| j                  |«      }t        ||| j                  ¬«      \  }}| j                  €¡t        | j                  t
        «      s!t        dt        | j                  «      › �«      ‚t        | j                  t
        «      s!t        dt        | j                  «      › �«      ‚| xj                  |z  c_        | xj                  |z  c_        yt        | j                  t        «      s!t        dt        | j                  «      › �«      ‚t        | j                  t        «      s!t        dt        | j                  «      › �«      ‚| j                  j                  |«       | j                  j                  |«       y)z*Update state with predictions and targets.N)r   z:Expected `self.sum_squared_error` to be a Tensor, but got z.Expected `self.total` to be a Tensor, but got z8Expected `self.sum_squared_error` to be a list, but got z,Expected `self.total` to be a list, but got )r0   r   r   r1   r#   r   Ú	TypeErrorÚtyper&   ÚlistÚappend)r6   r:   r;   r#   Únum_obss        r8   ÚupdatezPeakSignalNoiseRatio.updatep   sa  € à×ÑÐ'Ø×$Ñ$ UÓ+ˆEØ×%Ñ% fÓ-ˆFä%1°%¸ÀTÇXÁXÔ%NÑ"Ð˜7Ø�8‰8ÐÜ˜d×4Ñ4´fÔ=ÜØPÔQUÐVZ×VlÑVlÓQmÐPnÐoóð ô ˜dŸj™j¬&Ô1ÜÐ"PÔQUÐVZ×V`ÑV`ÓQaÐPbÐ cÓdÐdà×"Ò"Ð&7Ñ7Õ"Ø�JŠJ˜'Ñ!ŽJä˜d×4Ñ4´dÔ;ÜØNÌtÐTX×TjÑTjÓOkÐNlÐmóð ô ˜dŸj™j¬$Ô/ÜÐ"NÌtÐTX×T^ÑT^ÓO_ÐN`Ð aÓbÐbØ×"Ñ"×)Ñ)Ð*;Ô<Ø�J‰J×Ñ˜gÕ&r9   c                 óº  — t        | j                  t        j                  «      r| j                  }nat        | j                  t        «      r<t        j
                  | j                  D �cg c]  }|j                  «       ‘Œ c}«      }nt        d«      ‚t        | j                  t        j                  «      r| j                  }nat        | j                  t        «      r<t        j
                  | j                  D �cg c]  }|j                  «       ‘Œ c}«      }nt        d«      ‚t        ||| j                  | j                  | j                  ¬«      S c c}w c c}w )z.Compute peak signal-to-noise ratio over state.z>Expected sum_squared_error to be a Tensor or a list of Tensorsz2Expected total to be a Tensor or a list of Tensors)r   r   )r1   r#   r3   r   r?   r'   Úflattenr=   r&   r   r   r   r   )r6   r#   Úvaluer&   s       r8   ÚcomputezPeakSignalNoiseRatio.compute‹   së   € ä�d×,Ñ,¬e¯l©lÔ;Ø $× 6Ñ 6ÑÜ˜×.Ñ.´Ô5Ü %§	¡	È×H^ÑH^Ö*_¸u¨5¯=©=­?Ò*_Ó `ÑäÐ\Ó]Ð]ä�d—j‘j¤%§,¡,Ô/Ø—J‘J‰EÜ˜Ÿ
™
¤DÔ)Ü—I‘I¸D¿J¹JÖG°5˜uŸ}™}�ÒGÓH‰EäÐPÓQÐQäÐ.°°t·±ÈTÏYÉYÐbf×bpÑbpÔqÐqùò +`ùò Hs   Á)EÃ;EÚ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 PeakSignalNoiseRatio
            >>> metric = PeakSignalNoiseRatio(data_range=1.0)
            >>> preds = torch.tensor([[0.0, 1.0], [2.0, 3.0]])
            >>> target = torch.tensor([[3.0, 2.0], [1.0, 0.0]])
            >>> metric.update(preds, target)
            >>> fig_, ax_ = metric.plot()

        .. plot::
            :scale: 75

            >>> # Example plotting multiple values
            >>> import torch
            >>> from torchmetrics.image import PeakSignalNoiseRatio
            >>> metric = PeakSignalNoiseRatio(data_range=1.0)
            >>> preds = torch.tensor([[0.0, 1.0], [2.0, 3.0]])
            >>> target = torch.tensor([[3.0, 2.0], [1.0, 0.0]])
            >>> values = [ ]
            >>> for _ in range(10):
            ...     values.append(metric(preds, target))
            >>> fig_, ax_ = metric.plot(values)

        )Ú_plot)r6   rG   rH   s      r8   ÚplotzPeakSignalNoiseRatio.plot�   s   € ðX �z‰z˜#˜rÓ"Ð"r9   )g      $@r   N)NN)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   ÚboolÚ__annotations__r   r   r   r5   r   r   r2   r
   r   Úintr   r.   rB   rF   r   r   r   rK   Ú__classcell__)r7   s   @r8   r   r       s$  ø… ñ(ðT #Ð�tÓ"Ø!Ð�dÓ!Ø#Ð�tÓ#Ø!Ð�eÓ!ØÓð
 ØFXØ59ñDà˜%  u¨e |Ñ!4Ð4Ñ5ðDð ðDð ÐBÑCð	Dð
 �e˜C  s¨C x¡Ð0Ñ1Ñ2ðDð ðDð 
õDð>'˜Fð '¨Fð '°tó 'ð6r˜ó rð& _cñ,#Ø˜E &¨(°6Ñ*:Ð":Ñ;Ñ<ð,#ØIQÐRZÑI[ð,#à	÷,#r9   r   )Úcollections.abcr   Ú	functoolsr   Útypingr   r   r   r3   r   r	   Útyping_extensionsr
   Ú"torchmetrics.functional.image.psnrr   r   Útorchmetrics.metricr   Útorchmetrics.utilitiesr   Útorchmetrics.utilities.importsr   Útorchmetrics.utilities.plotr   r   Ú__doctest_skip__r   r,   r9   r8   ú<module>r^      sE   ðõ %Ý ß 'Ñ 'ã ß  Ý %ç JÝ &Ý 1Ý @ß @áØ3Ð4Ðôi#˜6õ i#r9   