Ë
    þÍ:j¦  ã                   ó  — d dl mZ d dlmZ d dl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 d	ed
edeeef   fd„Z	 	 	 dd	ed
edee   dee   deed      defd„Z	 	 	 dd	ed
edee   dee   deed      defd„Zy)é    )ÚSequence)ÚOptionalN)ÚTensorÚnn)ÚLiteral)Ú_gaussian_kernel_2d)Ú_check_same_shape)ÚreduceÚpredsÚtargetÚreturnc                 ó  — | j                   |j                   k7  r&t        d| j                   › d|j                   › d�«      ‚t        | |«       t        | j                  «      dk7  r&t        d| j                  › d|j                  › d�«      ‚| |fS )z¤Update and returns variables required to compute Universal Image Quality Index.

    Args:
        preds: Predicted tensor
        target: Ground truth tensor

    zEExpected `preds` and `target` to have the same data type. Got preds: z and target: ú.é   z@Expected `preds` and `target` to have BxCxHxW shape. Got preds: )ÚdtypeÚ	TypeErrorr	   ÚlenÚshapeÚ
ValueError)r   r   s     úv/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchmetrics/functional/image/uqi.pyÚ_uqi_updater      s—   € ð ‡{�{�f—l‘lÒ"ÜðØ Ÿ;™;˜- }°V·\±\°NÀ!ðEó
ð 	
ô �e˜VÔ$Ü
ˆ5�;‰;Ó˜1ÒÜØNÈuÏ{É{ÈmÐ[hÐio×iuÑiuÐhvÐvwÐxó
ð 	
ð �&ˆ=Ðó    Úkernel_sizeÚsigmaÚ	reduction)Úelementwise_meanÚsumÚnonec                 ó¬  — t        |«      dk7  st        |«      dk7  r$t        dt        |«      › dt        |«      › d�«      ‚t        d„ |D «       «      rt        d|› d�«      ‚t        d„ |D «       «      rt        d|› d�«      ‚| j                  }| j	                  d	«      }| j
                  }t        |||||«      }|d
   d	z
  dz  }	|d	   d	z
  dz  }
t        j                  j                  | |	|	|
|
fd¬«      } t        j                  j                  ||	|	|
|
fd¬«      }t        j                  | || | z  ||z  | |z  f«      }t        j                  j                  |||¬«      }|j                  | j                  d
   «      }|d
   j                  d«      }|d	   j                  d«      }|d
   |d	   z  }t        j                   |d   |z
  d¬«      }t        j                   |d   |z
  d¬«      }|d   |z
  }d|z  }||z   }t        j"                  |j
                  «      j$                  }d|z  |z  ||z   |z  |z   z  }|d|	|	 …|
|
 …f   }t'        ||«      S )a£  Compute Universal Image Quality Index.

    Args:
        preds: estimated image
        target: ground truth image
        kernel_size: size of the gaussian kernel
        sigma: Standard deviation of the gaussian kernel
        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.rand([16, 1, 16, 16])
        >>> target = preds * 0.75
        >>> preds, target = _uqi_update(preds, target)
        >>> _uqi_compute(preds, target)
        tensor(0.9216)

    é   zOExpected `kernel_size` and `sigma` to have the length of two. Got kernel_size: z and sigma: r   c              3   ó:   K  — | ]  }|d z  dk(  xs |dk  –— Œ y­w)r    r   N© )Ú.0Úxs     r   ú	<genexpr>z_uqi_compute.<locals>.<genexpr>Q   s$   è ø€ Ò
5 Aˆ1ˆq‰5�A‰:Ò˜˜a™ÓÑ
5ùs   ‚z8Expected `kernel_size` to have odd positive number. Got c              3   ó&   K  — | ]	  }|d k  –— Œ y­w)r   Nr"   )r#   Úys     r   r%   z_uqi_compute.<locals>.<genexpr>T   s   è ø€ Ò
!�aˆ1��6Ñ
!ùs   ‚z.Expected `sigma` to have positive number. Got é   r   Úreflect)Úmode)Úgroupsg        )Úminé   r   .)r   r   ÚanyÚdeviceÚsizer   r   r   Ú
functionalÚpadÚtorchÚcatÚconv2dÚsplitr   ÚpowÚclampÚfinfoÚepsr
   )r   r   r   r   r   r/   Úchannelr   ÚkernelÚpad_hÚpad_wÚ
input_listÚoutputsÚoutput_listÚ
mu_pred_sqÚmu_target_sqÚmu_pred_targetÚsigma_pred_sqÚsigma_target_sqÚsigma_pred_targetÚupperÚlowerr:   Úuqi_idxs                           r   Ú_uqi_computerK   /   s†  € ô8 ˆ;Ó˜1Ò¤ E£
¨a¢Üð!Ü!$ [Ó!1Ð 2°,¼sÀ5»z¸lÈ!ðMó
ð 	
ô
 Ñ
5¨Ô
5Ô5ÜÐSÐT_ÐS`Ð`aÐbÓcÐcä
Ñ
!˜5Ô
!Ô!ÜÐIÈ%ÈÐPQÐRÓSÐSà�\‰\€FØ�j‰j˜‹m€GØ�K‰K€EÜ  ¨+°u¸eÀVÓL€FØ˜‰^˜aÑ AÑ%€EØ˜‰^˜aÑ AÑ%€Eä�M‰M×Ñ˜e e¨U°E¸5Ð%AÈ	ÐÓR€EÜ�]‰]×Ñ˜v¨¨u°e¸UÐ'CÈ)ÐÓT€Fä—‘˜E 6¨5°5©=¸&À6¹/È5ÐSYÉ>ÐZÓ[€JÜ�m‰m×"Ñ" :¨v¸gÐ"ÓF€GØ—-‘- §¡¨A¡Ó/€Kà˜Q‘×#Ñ# AÓ&€JØ˜q‘>×%Ñ% aÓ(€LØ  ‘^ k°!¡nÑ4€Nô —K‘K ¨A¡°Ñ ;ÀÔE€MÜ—k‘k +¨a¡.°<Ñ"?ÀSÔI€OØ# A™¨Ñ7ÐàÐ!Ñ!€EØ˜OÑ+€EÜ
�+‰+�m×)Ñ)Ó
*×
.Ñ
.€CØ�NÑ" eÑ+°¸lÑ1JÈeÑ0SÐVYÑ0YÑZ€GØ�c˜5 % ˜<¨°¨v¨Ð5Ñ6€Gä�'˜9Ó%Ð%r   c                 ó>   — t        | |«      \  } }t        | ||||«      S )aÐ  Universal Image Quality Index.

    Args:
        preds: estimated image
        target: ground truth image
        kernel_size: size of the gaussian kernel
        sigma: Standard deviation of the gaussian kernel
        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

    Return:
        Tensor with UniversalImageQualityIndex score

    Raises:
        TypeError:
            If ``preds`` and ``target`` don't have the same data type.
        ValueError:
            If ``preds`` and ``target`` don't have ``BxCxHxW shape``.
        ValueError:
            If the length of ``kernel_size`` or ``sigma`` is not ``2``.
        ValueError:
            If one of the elements of ``kernel_size`` is not an ``odd positive number``.
        ValueError:
            If one of the elements of ``sigma`` is not a ``positive number``.

    Example:
        >>> from torchmetrics.functional.image import universal_image_quality_index
        >>> preds = torch.rand([16, 1, 16, 16])
        >>> target = preds * 0.75
        >>> universal_image_quality_index(preds, target)
        tensor(0.9216)

    References:
        [1] Zhou Wang and A. C. Bovik, "A universal image quality index," in IEEE Signal Processing Letters, vol. 9,
        no. 3, pp. 81-84, March 2002, doi: 10.1109/97.995823.

        [2] Zhou Wang, A. C. Bovik, H. R. Sheikh and E. P. Simoncelli, "Image quality assessment: from error visibility
        to structural similarity," in IEEE Transactions on Image Processing, vol. 13, no. 4, pp. 600-612, April 2004,
        doi: 10.1109/TIP.2003.819861.

    )r   rK   )r   r   r   r   r   s        r   Úuniversal_image_quality_indexrM   w   s(   € ôf    vÓ.�M€Eˆ6Ü˜˜v {°E¸9ÓEÐEr   ))é   rN   )ç      ø?rO   r   )Úcollections.abcr   Útypingr   r3   r   r   Útyping_extensionsr   Ú#torchmetrics.functional.image.utilsr   Útorchmetrics.utilities.checksr	   Ú"torchmetrics.utilities.distributedr
   Útupler   ÚintÚfloatrK   rM   r"   r   r   ú<module>rY      s  ðõ %Ý ã ß Ý %å CÝ ;Ý 5ð�vð  vð °%¸À¸Ñ2Gó ð0 "*Ø'ØFXñE&ØðE&àðE&ð ˜#‘ðE&ð �E‰?ð	E&ð
 ˜Ð AÑBÑCðE&ð óE&ðV "*Ø'ØFXñ4FØð4Fàð4Fð ˜#‘ð4Fð �E‰?ð	4Fð
 ˜Ð AÑBÑCð4Fð ô4Fr   