Ë
    þÍ:j¸  ã                   ó  — d dl 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ededed	ed
ed   defd„Z	 ddededeeeeedf   f      deeef   fd„Z	 	 	 ddededeeeeef   f   d	ed
ed   deeeeedf   f      defd„Zy)é    )ÚOptionalÚUnionN)ÚTensorÚtensor)ÚLiteral)Úrank_zero_warnÚreduceÚsum_squared_errorÚnum_obsÚ
data_rangeÚbaseÚ	reduction)Úelementwise_meanÚsumÚnoneNÚreturnc                 óÆ   — dt        j                  |«      z  t        j                  | |z  «      z
  }|dt        j                  t        |«      «      z  z  }t        ||¬«      S )aÃ  Compute peak signal-to-noise ratio.

    Args:
        sum_squared_error: Sum of square of errors over all observations
        num_obs: Number of predictions or observations
        data_range: the range of the data. If None, it is determined from the data (max - min).
           ``data_range`` must be given when ``dim`` is not None.
        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

    Example:
        >>> preds = torch.tensor([[0.0, 1.0], [2.0, 3.0]])
        >>> target = torch.tensor([[3.0, 2.0], [1.0, 0.0]])
        >>> data_range = target.max() - target.min()
        >>> sum_squared_error, num_obs = _psnr_update(preds, target)
        >>> _psnr_compute(sum_squared_error, num_obs, data_range)
        tensor(2.5527)

    é   é
   )r   )ÚtorchÚlogr   r	   )r
   r   r   r   r   Úpsnr_base_eÚ	psnr_valss          úw/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchmetrics/functional/image/psnr.pyÚ_psnr_computer      sT   € ð< ”e—i‘i 
Ó+Ñ+¬e¯i©iÐ8IÈGÑ8SÓ.TÑT€KØ˜r¤E§I¡I¬f°T«lÓ$;Ñ;Ñ<€IÜ�) yÔ1Ð1ó    ÚpredsÚtargetÚdim.c                 óÒ  — | j                  «       s| j                  t        j                  «      } |j                  «       s|j                  t        j                  «      }|€Ut        j                  t        j
                  | |z
  d«      «      }t        |j                  «       |j                  ¬«      }||fS | |z
  }t        j                  ||z  |¬«      }t        |t        «      r|gn
t        |«      }|s)t        |j                  «       |j                  ¬«      }||fS t        |j                  «       |j                  ¬«      |   j                  «       }|j                  |«      }||fS )aZ  Update and return variables required to compute peak signal-to-noise ratio.

    Args:
        preds: Predicted tensor
        target: Ground truth tensor
        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.

    r   )Údevice©r   )Úis_floating_pointÚtor   Úfloat32r   Úpowr   Únumelr!   Ú
isinstanceÚintÚlistÚsizeÚprodÚ	expand_as)r   r   r   r
   r   ÚdiffÚdim_lists          r   Ú_psnr_updater0   :   s   € ð ×"Ñ"Ô$Ø—‘œŸ™Ó'ˆØ×#Ñ#Ô%Ø—‘œ5Ÿ=™=Ó)ˆà
€{Ü!ŸI™I¤e§i¡i°¸±ÀÓ&BÓCÐÜ˜Ÿ™›°·±Ô>ˆØ  'Ð)Ð)à�6‰>€DÜŸ	™	 $¨¡+°3Ô7Ðä" 3¬Ô,�‰u´$°s³)€HÙÜ˜Ÿ™›°·±Ô>ˆð
 ˜gÐ%Ð%ô ˜Ÿ™›¨v¯}©}Ô=¸hÑG×LÑLÓNˆØ×#Ñ#Ð$5Ó6ˆà˜gÐ%Ð%r   c                 óZ  — |€|dk7  rt        d|› d�«       t        |t        «      rQt        j                  | |d   |d   ¬«      } t        j                  ||d   |d   ¬«      }t        |d   |d   z
  «      }nt        t        |«      «      }t        | ||¬«      \  }}t        |||||¬«      S )	a±  Compute the peak signal-to-noise ratio.

    Args:
        preds: estimated signal
        target: groun truth signal
        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.

    Return:
        Tensor with PSNR score

    Example:
        >>> from torchmetrics.functional.image import peak_signal_noise_ratio
        >>> pred = torch.tensor([[0.0, 1.0], [2.0, 3.0]])
        >>> target = torch.tensor([[3.0, 2.0], [1.0, 0.0]])
        >>> peak_signal_noise_ratio(pred, target, data_range=3.0)
        tensor(2.5527)

    .. attention::
        Half precision is only support on GPU for this metric.

    r   zThe `reduction=z.` will not have any effect when `dim` is None.r   é   )ÚminÚmaxr"   )r   r   )	r   r(   Útupler   Úclampr   Úfloatr0   r   )	r   r   r   r   r   r   Údata_range_valr
   r   s	            r   Úpeak_signal_noise_ratior9   _   s±   € ðR €{�yÐ$6Ò6Ü˜¨¨Ð3aÐbÔcä�*œeÔ$Ü—‘˜E z°!¡}¸*ÀQ¹-ÔHˆÜ—‘˜V¨°A©¸JÀq¹MÔJˆÜ 
¨1¡°
¸1±Ñ =Ó>‰ä¤ jÓ 1Ó2ˆä!-¨e°VÀÔ!EÑÐ�wÜÐ*¨G°^È$ÐZcÔdÐdr   )ç      $@r   )N)r:   r   N)Útypingr   r   r   r   r   Útyping_extensionsr   Útorchmetrics.utilitiesr   r	   r7   r   r)   r5   r0   r9   © r   r   ú<module>r?      s;  ð÷ #ã ß  Ý %ç 9ð ØBTñ 2Øð 2àð 2ð ð 2ð ð	 2ð
 Ð>Ñ?ð 2ð ó 2ðL 26ñ"&Øð"&àð"&ð 
�%˜˜U 3¨ 8™_Ð,Ñ-Ñ	.ð"&ð ˆ6�6ˆ>Ñó	"&ðR ØBTØ15ñ4eØð4eàð4eð �e˜U 5¨% <Ñ0Ð0Ñ1ð4eð ð	4eð
 Ð>Ñ?ð4eð 
�%˜˜U 3¨ 8™_Ð,Ñ-Ñ	.ð4eð ô4er   