Ë
    þÍ:j»  ã                   ó´   — d dl mZ d dlmZmZmZm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 d d
lmZmZ esdgZ G d„ de«      Zy)é    )ÚSequence)ÚAnyÚListÚOptionalÚUnion)ÚTensorÚtensor)ÚLiteral)Ú_uqi_computeÚ_uqi_update)ÚMetric)Úrank_zero_warn)Údim_zero_cat)Ú_MATPLOTLIB_AVAILABLE)Ú_AX_TYPEÚ_PLOT_OUT_TYPEzUniversalImageQualityIndex.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<   d	Ze
ed
<   ee   ed<   ee   ed<   eed<   eed<   	 	 	 ddee   dee
   ded   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 )ÚUniversalImageQualityIndexa…  Compute Universal Image Quality Index (UniversalImageQualityIndex_).

    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

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

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

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

    Return:
        Tensor with UniversalImageQualityIndex score

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

    TÚis_differentiableÚhigher_is_betterFÚfull_state_updateç        Úplot_lower_boundg      ð?Úplot_upper_boundÚpredsÚtargetÚsum_uqiÚnumelÚkernel_sizeÚsigmaÚ	reduction©Úelementwise_meanÚsumÚnoneNÚkwargsÚreturnNc                 ó^  •— t        ‰| �  di |¤Ž |dvrt        d|› d�«      ‚|�|dk(  r4t        d«       | j	                  dg d¬«       | j	                  d	g d¬«       n:| j	                  d
t        d«      d¬«       | j	                  dt        d«      d¬«       || _        || _        || _        y )Nr"   zThe `reduction` zI is not valid. Valid options are `elementwise_mean`, `sum`, `none`, None.r%   zÇMetric `UniversalImageQualityIndex` will save all targets and predictions in the buffer when using`reduction=None` or `reduction='none'. For large datasets, this may lead to a large memory footprint.r   Úcat)ÚdefaultÚdist_reduce_fxr   r   r   r$   )r+   r   r   © )	ÚsuperÚ__init__Ú
ValueErrorr   Ú	add_stater	   r   r    r!   )Úselfr   r    r!   r&   Ú	__class__s        €úk/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchmetrics/image/uqi.pyr.   z#UniversalImageQualityIndex.__init__P   s½   ø€ ô 	‰ÑÑ"˜6Ò"ØÐEÑEÜØ" 9 +Ð-vÐwóð ð Ð 	¨VÒ 3Üðxôð �N‰N˜7¨B¸uˆNÔEØ�N‰N˜8¨RÀˆNÕFà�N‰N˜9¤f¨S£kÀ%ˆNÔHØ�N‰N˜7¤F¨1£I¸eˆNÔDØ&ˆÔØˆŒ
Ø"ˆ�ó    c                 óð  — t        ||«      \  }}| j                  �| j                  dk(  r7| j                  j                  |«       | j                  j                  |«       yt        ||| j                  | j                  d¬«      }| xj                  |z  c_        |j                  }| xj                  |d   |d   z  |d   | j                  d   z
  dz   z  |d   | j                  d   z
  dz   z  z  c_
        y)	z*Update state with predictions and targets.Nr%   r$   )r!   r   é   é   é   )r   r!   r   Úappendr   r   r   r    r   Úshaper   )r1   r   r   Ú	uqi_scoreÚpss        r3   Úupdatez!UniversalImageQualityIndex.updatej   sÙ   € ä# E¨6Ó2‰ˆˆvØ�>‰>Ð! T§^¡^°vÒ%=Ø�J‰J×Ñ˜eÔ$Ø�K‰K×Ñ˜vÕ&ä$ U¨F°D×4DÑ4DÀdÇjÁjÐ\aÔbˆIØ�LŠL˜IÑ%�LØ—‘ˆBØ�JŠJ˜"˜Q™% " Q¡%™-¨2¨a©5°4×3CÑ3CÀAÑ3FÑ+FÈÑ+JÑKÈrÐRSÉuÐW[×WgÑWgÐhiÑWjÑOjÐmnÑOnÑoÑoŽJr4   c                 óN  — | j                   dk(  s| j                   €Wt        | j                  «      }t        | j                  «      }t	        ||| j
                  | j                  | j                   «      S | j                   dk(  r| j                  | j                  z  S | j                  S )z&Compute explained variance over state.r%   r#   )	r!   r   r   r   r   r   r    r   r   )r1   r   r   s      r3   Úcomputez"UniversalImageQualityIndex.computev   s   € à�>‰>˜VÒ# t§~¡~Ð'=Ü  §¡Ó,ˆEÜ! $§+¡+Ó.ˆFÜ  v¨t×/?Ñ/?ÀÇÁÈTÏ^É^Ó\Ð\Ø,0¯N©NÐ>PÒ,Pˆt�|‰|˜dŸj™jÑ(ÐbÐVZ×VbÑVbÐbr4   Ú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 UniversalImageQualityIndex
            >>> preds = torch.rand([16, 1, 16, 16])
            >>> target = preds * 0.75
            >>> metric = UniversalImageQualityIndex()
            >>> metric.update(preds, target)
            >>> fig_, ax_ = metric.plot()

        .. plot::
            :scale: 75

            >>> # Example plotting multiple values
            >>> import torch
            >>> from torchmetrics.image import UniversalImageQualityIndex
            >>> preds = torch.rand([16, 1, 16, 16])
            >>> target = preds * 0.75
            >>> metric = UniversalImageQualityIndex()
            >>> values = [ ]
            >>> for _ in range(10):
            ...     values.append(metric(preds, target))
            >>> fig_, ax_ = metric.plot(values)

        )Ú_plot)r1   r@   rA   s      r3   ÚplotzUniversalImageQualityIndex.plot~   s   € ðX �z‰z˜#˜rÓ"Ð"r4   ))é   rE   )ç      ø?rF   r#   )NN)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   ÚboolÚ__annotations__r   r   r   Úfloatr   r   r   r   Úintr
   r   r.   r=   r?   r   r   r   r   rD   Ú__classcell__)r2   s   @r3   r   r      s  ø… ñ#ðJ #Ð�tÓ"Ø!Ð�dÓ!Ø#Ð�tÓ#Ø!Ð�eÓ!Ø!Ð�eÓ!à�‰<ÓØ�‰LÓØƒOØƒMð &.Ø!+ØFXñ	#à˜c‘]ð#ð ˜‰ð#ð ÐBÑCð	#ð
 ð#ð 
õ#ð4
p˜Fð 
p¨Fð 
p°tó 
pðc˜ó cð _cñ,#Ø˜E &¨(°6Ñ*:Ð":Ñ;Ñ<ð,#ØIQÐRZÑI[ð,#à	÷,#r4   r   N)Úcollections.abcr   Útypingr   r   r   r   Útorchr   r	   Útyping_extensionsr
   Ú!torchmetrics.functional.image.uqir   r   Útorchmetrics.metricr   Útorchmetrics.utilitiesr   Útorchmetrics.utilities.datar   Útorchmetrics.utilities.importsr   Útorchmetrics.utilities.plotr   r   Ú__doctest_skip__r   r,   r4   r3   ú<module>r[      sB   ðõ %ß -Ó -ç  Ý %ç GÝ &Ý 1Ý 4Ý @ß @áØ9Ð:ÐôK# õ K#r4   