Ë
    þÍ:jv  ã                   ó°   — d dl mZ d dlmZmZmZmZ d dl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)ÚLiteral)Ú"_spectral_distortion_index_computeÚ!_spectral_distortion_index_update)ÚMetric)Úrank_zero_warn)Údim_zero_cat)Ú_MATPLOTLIB_AVAILABLE)Ú_AX_TYPEÚ_PLOT_OUT_TYPEzSpectralDistortionIndex.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<   	 dd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 )ÚSpectralDistortionIndexag  Compute Spectral Distortion Index (SpectralDistortionIndex_) also now as D_lambda.

    The metric is used to compare the spectral distortion between two images.

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

    - ``preds`` (:class:`~torch.Tensor`): Low resolution multispectral image of shape ``(N,C,H,W)``
    - ``target``(:class:`~torch.Tensor`): High resolution fused image of shape ``(N,C,H,W)``

    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 SDI value
      over sample else returns tensor of shape ``(N,)`` with SDI values per sample

    Args:
        p: Large spectral differences
        reduction: a method to reduce metric score over labels.

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

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

    Example:
        >>> from torch import rand
        >>> from torchmetrics.image import SpectralDistortionIndex
        >>> preds = rand([16, 3, 16, 16])
        >>> target = rand([16, 3, 16, 16])
        >>> sdi = SpectralDistortionIndex()
        >>> sdi(preds, target)
        tensor(0.0234)

    TÚhigher_is_betterÚis_differentiableFÚfull_state_updateg        Úplot_lower_boundg      ð?Úplot_upper_boundÚpredsÚtargetÚpÚ	reduction©Úelementwise_meanÚsumÚnoneÚkwargsÚreturnNc                 ó  •— t        ‰| �  di |¤Ž t        d«       t        |t        «      r|dk  rt        d|› d�«      ‚|| _        d}||vrt        d|› d|› �«      ‚|| _        | j                  dg d	¬
«       | j                  dg d	¬
«       y )Nz�Metric `SpectralDistortionIndex` will save all targets and predictions in buffer. For large datasets this may lead to large memory footprint.r   z.Expected `p` to be a positive integer. Got p: ú.r   z(Expected argument `reduction` be one of z	 but got r   Úcat)ÚdefaultÚdist_reduce_fxr   © )	ÚsuperÚ__init__r   Ú
isinstanceÚintÚ
ValueErrorr   r   Ú	add_state)Úselfr   r   r!   Úallowed_reductionsÚ	__class__s        €úp/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchmetrics/image/d_lambda.pyr*   z SpectralDistortionIndex.__init__L   s©   ø€ ô 	‰ÑÑ"˜6Ò"Üð*ô	
ô ˜!œSÔ! Q¨!¢VÜÐMÈaÈSÐPQÐRÓSÐSØˆŒØ@ÐØÐ.Ñ.ÜÐGÐHZÐG[Ð[dÐenÐdoÐpÓqÐqØ"ˆŒØ�‰�w¨¸5ˆÔAØ�‰�x¨¸EˆÕBó    c                 óŽ   — t        ||«      \  }}| j                  j                  |«       | j                  j                  |«       y)z#Update state with preds and target.N)r   r   Úappendr   ©r/   r   r   s      r2   ÚupdatezSpectralDistortionIndex.update`   s6   € ä9¸%ÀÓH‰ˆˆvØ�
‰
×Ñ˜%Ô Ø�‰×Ñ˜6Õ"r3   c                 óš   — t        | j                  «      }t        | j                  «      }t        ||| j                  | j
                  «      S )z.Compute and returns spectral distortion index.)r   r   r   r
   r   r   r6   s      r2   ÚcomputezSpectralDistortionIndex.computef   s7   € ä˜TŸZ™ZÓ(ˆÜ˜dŸk™kÓ*ˆÜ1°%¸ÀÇÁÈÏÉÓXÐXr3   Ú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
            >>> from torch import rand
            >>> from torchmetrics.image import SpectralDistortionIndex
            >>> preds = rand([16, 3, 16, 16])
            >>> target = rand([16, 3, 16, 16])
            >>> metric = SpectralDistortionIndex()
            >>> metric.update(preds, target)
            >>> fig_, ax_ = metric.plot()

        .. plot::
            :scale: 75

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

        )Ú_plot)r/   r:   r;   s      r2   ÚplotzSpectralDistortionIndex.plotl   s   € ðX �z‰z˜#˜rÓ"Ð"r3   )é   r   )NN)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   ÚboolÚ__annotations__r   r   r   Úfloatr   r   r   r,   r	   r   r*   r7   r9   r   r   r   r   r   r>   Ú__classcell__)r1   s   @r2   r   r      só   ø… ñ!ðF "Ð�dÓ!Ø"Ð�tÓ"Ø#Ð�tÓ#Ø!Ð�eÓ!Ø!Ð�eÓ!à�‰<ÓØ�‰LÓð SeñCØðCØ%,Ð-NÑ%OðCØpsðCà	õCð(#˜Fð #¨Fð #°tó #ðY˜ó Yð _cñ,#Ø˜E &¨(°6Ñ*:Ð":Ñ;Ñ<ð,#ØIQÐRZÑI[ð,#à	÷,#r3   r   N)Úcollections.abcr   Útypingr   r   r   r   Útorchr   Útyping_extensionsr	   Ú&torchmetrics.functional.image.d_lambdar
   r   Útorchmetrics.metricr   Útorchmetrics.utilitiesr   Útorchmetrics.utilities.datar   Útorchmetrics.utilities.importsr   Útorchmetrics.utilities.plotr   r   Ú__doctest_skip__r   r(   r3   r2   ú<module>rS      sB   ðõ %ß -Ó -å Ý %ç xÝ &Ý 1Ý 4Ý @ß @áØ6Ð7Ðôy#˜fõ y#r3   