Ë
    þÍ:jš%  ã                   óÐ   — 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mZ d dlmZ d dlmZ d d	lmZ d d
lmZmZ d dlmZmZ esdgZesddgZ G d„ de«      Zy)é    )ÚSequence)ÚAnyÚListÚOptionalÚUnion)ÚTensor)ÚLiteral)Ú"_spectral_distortion_index_computeÚ!_spectral_distortion_index_update)Ú!_spatial_distortion_index_computeÚ _spatial_distortion_index_update)ÚMetric)Úrank_zero_warn)Údim_zero_cat)Ú_MATPLOTLIB_AVAILABLEÚ_TORCHVISION_AVAILABLE)Ú_AX_TYPEÚ_PLOT_OUT_TYPEzQualityWithNoReference.plotÚQualityWithNoReferencec                   ó8  ‡ — 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   ed<   ee   ed<   	 	 	 	 	 dde
de
dededed   deddfˆ fd„Zdedeeef   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 )!r   aì  Compute Quality with No Reference (QualityWithNoReference_) also now as QNR.

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

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

    - ``preds`` (:class:`~torch.Tensor`): High resolution multispectral image of shape ``(N,C,H,W)``.
    - ``target`` (:class:`~Dict`): A dictionary containing the following keys:

      - ``ms`` (:class:`~torch.Tensor`): Low resolution multispectral image of shape ``(N,C,H',W')``.
      - ``pan`` (:class:`~torch.Tensor`): High resolution panchromatic image of shape ``(N,C,H,W)``.
      - ``pan_lr`` (:class:`~torch.Tensor`): (optional) Low resolution panchromatic image of shape ``(N,C,H',W')``.

    where H and W must be multiple of H' and W'.

    When ``pan_lr`` is ``None``, a uniform filter will be applied on ``pan`` to produce a degraded image. The degraded
    image is then resized to match the size of ``ms`` and served as ``pan_lr`` in the calculation.

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

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

    Args:
        alpha: Relevance of spectral distortion.
        beta: Relevance of spatial distortion.
        norm_order: Order of the norm applied on the difference.
        window_size: Window size of the filter applied to degrade the high resolution panchromatic image.
        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 QualityWithNoReference
        >>> preds = rand([16, 3, 32, 32])
        >>> target = {
        ...     'ms': rand([16, 3, 16, 16]),
        ...     'pan': rand([16, 3, 32, 32]),
        ... }
        >>> qnr = QualityWithNoReference()
        >>> qnr(preds, target)
        tensor(0.9694)

    TÚhigher_is_betterÚis_differentiableFÚfull_state_updateg        Úplot_lower_boundg      ð?Úplot_upper_boundÚpredsÚmsÚpanÚpan_lrÚalphaÚbetaÚ
norm_orderÚwindow_sizeÚ	reduction©Úelementwise_meanÚsumÚnoneÚkwargsÚreturnNc                 ó†  •— t        ‰| �  di |¤Ž t        d«       t        |t        t
        f«      r|dk  rt        d|› d�«      ‚|| _        t        |t        t
        f«      r|dk  rt        d|› d�«      ‚|| _        t        |t        «      r|dk  rt        d|› d�«      ‚|| _	        t        |t        «      r|dk  rt        d|› d�«      ‚|| _
        d}||vrt        d	|› d
|› �«      ‚|| _        | j                  dg d¬«       | j                  dg d¬«       | j                  dg d¬«       | j                  dg d¬«       y )NzŒMetric `QualityWithNoReference` will save all targets and predictions in buffer. For large datasets this may lead to large memory footprint.r   z>Expected `alpha` to be a non-negative real number. Got alpha: ú.z<Expected `beta` to be a non-negative real number. Got beta: z@Expected `norm_order` to be a positive integer. Got norm_order: zBExpected `window_size` to be a positive integer. Got window_size: r%   z(Expected argument `reduction` be one of z	 but got r   Úcat)ÚdefaultÚdist_reduce_fxr   r   r   © )ÚsuperÚ__init__r   Ú
isinstanceÚintÚfloatÚ
ValueErrorr    r!   r"   r#   r$   Ú	add_state)	Úselfr    r!   r"   r#   r$   r)   Úallowed_reductionsÚ	__class__s	           €úk/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchmetrics/image/qnr.pyr2   zQualityWithNoReference.__init__b   sg  ø€ ô 	‰ÑÑ"˜6Ò"ÜðKô	
ô
 ˜%¤#¤u Ô.°%¸!²)ÜÐ]Ð^cÐ]dÐdeÐfÓgÐgØˆŒ
Ü˜$¤¤e Ô-°¸²ÜÐ[Ð\`Ð[aÐabÐcÓdÐdØˆŒ	Ü˜*¤cÔ*¨j¸AªoÜÐ_Ð`jÐ_kÐklÐmÓnÐnØ$ˆŒÜ˜+¤sÔ+¨{¸aÒ/?ÜÐaÐbmÐanÐnoÐpÓqÐqØ&ˆÔØ@ÐØÐ.Ñ.ÜÐGÐHZÐG[Ð[dÐenÐdoÐpÓqÐqØ"ˆŒØ�‰�w¨¸5ˆÔAØ�‰�t R¸ˆÔ>Ø�‰�u b¸ˆÔ?Ø�‰�x¨¸EˆÕBó    Útargetc                 óö  — d|vrt        d|j                  «       › d�«      ‚d|vrt        d|j                  «       › d�«      ‚|d   }|d   }|j                  d«      }t        ||«      \  }}t	        ||||«      \  }}}}| j
                  j                  |«       | j                  j                  |d   «       | j                  j                  |d   «       d|v r| j                  j                  |d   «       yy)aì  Update state with preds and target.

        Args:
            preds: High resolution multispectral image.
            target: A dictionary containing the following keys:

                - ``'ms'``: low resolution multispectral image.
                - ``'pan'``: high resolution panchromatic image.
                - ``'pan_lr'``: (optional) low resolution panchromatic image.

        Raises:
            ValueError:
                If ``target`` doesn't have ``ms`` and ``pan``.

        r   z0Expected `target` to have key `ms`. Got target: r,   r   z1Expected `target` to have key `pan`. Got target: r   N)
r6   ÚkeysÚgetr   r   r   Úappendr   r   r   )r8   r   r=   r   r   r   s         r;   ÚupdatezQualityWithNoReference.update†   sý   € ð  �vÑÜÐOÐPV×P[ÑP[ÓP]ÈÐ^_Ð`ÓaÐaØ˜ÑÜÐPÐQW×Q\ÑQ\ÓQ^ÐP_Ð_`ÐaÓbÐbØ�D‰\ˆØ�U‰mˆØ—‘˜HÓ%ˆÜ5°e¸RÓ@‰	ˆˆrÜ!AÀ%ÈÈSÐRXÓ!YÑˆˆr�3˜Ø�
‰
×Ñ˜%Ô Ø�‰�‰�v˜d‘|Ô$Ø�‰�‰˜˜u™Ô&Ø�vÑØ�K‰K×Ñ˜v hÑ/Õ0ð r<   c           	      óÊ  — t        | j                  «      }t        | j                  «      }t        | j                  «      }t	        | j
                  «      dkD  rt        | j
                  «      nd}t        ||| j                  | j                  «      }t        ||||| j                  | j                  | j                  «      }d|z
  | j                  z  d|z
  | j                  z  z  S )z.Compute and returns quality with no reference.r   Né   )r   r   r   r   Úlenr   r
   r"   r$   r   r#   r    r!   )r8   r   r   r   r   Úd_lambdaÚd_ss          r;   ÚcomputezQualityWithNoReference.compute¥   sµ   € ä˜TŸZ™ZÓ(ˆÜ˜$Ÿ'™'Ó"ˆÜ˜4Ÿ8™8Ó$ˆÜ.1°$·+±+Ó.>ÀÒ.B”˜dŸk™kÔ*ÈˆÜ5°e¸RÀÇÁÐRV×R`ÑR`ÓaˆÜ/Ø�2�s˜F D§O¡O°T×5EÑ5EÀtÇ~Á~ó
ˆð �H‘ §¡Ñ+¨q°3©w¸4¿9¹9Ñ.DÑDÐDr<   Ú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 QualityWithNoReference
            >>> preds = rand([16, 3, 32, 32])
            >>> target = {
            ...     'ms': rand([16, 3, 16, 16]),
            ...     'pan': rand([16, 3, 32, 32]),
            ... }
            >>> metric = QualityWithNoReference()
            >>> metric.update(preds, target)
            >>> fig_, ax_ = metric.plot()

        .. plot::
            :scale: 75

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

        )Ú_plot)r8   rI   rJ   s      r;   ÚplotzQualityWithNoReference.plot±   s   € ðd �z‰z˜#˜rÓ"Ð"r<   )rD   rD   rD   é   r&   )NN)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   ÚboolÚ__annotations__r   r   r   r5   r   r   r   r4   r	   r   r2   ÚdictÚstrrB   rH   r   r   r   r   r   rM   Ú__classcell__)r:   s   @r;   r   r   $   sB  ø… ñ0ðd "Ð�dÓ!Ø"Ð�tÓ"Ø#Ð�tÓ#Ø!Ð�eÓ!Ø!Ð�eÓ!à�‰<ÓØˆV‰ÓØ	ˆf‰ÓØ�‰LÓð ØØØØ@Rñ"Càð"Cð ð"Cð ð	"Cð
 ð"Cð Ð<Ñ=ð"Cð ð"Cð 
õ"CðH1˜Fð 1¨D°°f°Ñ,=ð 1À$ó 1ð>
E˜ó 
Eð _cñ2#Ø˜E &¨(°6Ñ*:Ð":Ñ;Ñ<ð2#ØIQÐRZÑI[ð2#à	÷2#r<   N)Úcollections.abcr   Útypingr   r   r   r   Útorchr   Útyping_extensionsr	   Ú&torchmetrics.functional.image.d_lambdar
   r   Ú!torchmetrics.functional.image.d_sr   r   Útorchmetrics.metricr   Útorchmetrics.utilitiesr   Útorchmetrics.utilities.datar   Útorchmetrics.utilities.importsr   r   Útorchmetrics.utilities.plotr   r   Ú__doctest_skip__r   r0   r<   r;   ú<module>rd      sT   ðõ %ß -Ó -å Ý %ç xß qÝ &Ý 1Ý 4ß Xß @áØ5Ð6ÐáØ0Ð2OÐPÐô#˜Võ #r<   