Ë
    þÍ:jn  ã                   ón   — d dl mZmZ d dl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  G d„ d	e
«      Zy)
é    )ÚAnyÚListN)ÚTensor)ÚLiteral)Ú_vif_per_channel)ÚMetric)Údim_zero_catc            	       ó„   ‡ — e Zd ZU dZdZdZdZee   e	d<   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ˆ xZS )ÚVisualInformationFidelityu<  Compute Pixel Based Visual Information Fidelity (VIF_).

    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)`` with H,W â‰¥ 41
    - ``target`` (:class:`~torch.Tensor`): Ground truth values of shape ``(N,C,H,W)`` with H,W â‰¥ 41

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

    - ``vif-p`` (:class:`~torch.Tensor`):
        - If ``reduction='mean'`` (default), returns a Tensor mean VIF score.
        - If ``reduction='none'``, returns a tensor of shape ``(N,)`` with VIF values per sample.

    Args:
        sigma_n_sq: variance of the visual noise
        reduction: The reduction method for aggregating scores.

            - ``'mean'``: return the average VIF across the batch.
            - ``'none'``: return a VIF score for each sample in the batch.

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

    Example:
        >>> from torch import randn
        >>> from torchmetrics.image import VisualInformationFidelity
        >>> preds = randn([32, 3, 41, 41], generator=torch.Generator().manual_seed(42))
        >>> target = randn([32, 3, 41, 41], generator=torch.Generator().manual_seed(43))
        >>> vif_mean = VisualInformationFidelity(reduction='mean')
        >>> vif_mean(preds, target)
        tensor(0.0032)
        >>> vif_none = VisualInformationFidelity(reduction='none')
        >>> vif_none(preds, target)
        tensor([0.0040, 0.0049, 0.0017, 0.0039, 0.0041, 0.0043, 0.0030, 0.0028, 0.0012,
                0.0067, 0.0010, 0.0014, 0.0030, 0.0048, 0.0050, 0.0038, 0.0037, 0.0025,
                0.0041, 0.0019, 0.0007, 0.0034, 0.0037, 0.0016, 0.0026, 0.0021, 0.0038,
                0.0033, 0.0031, 0.0020, 0.0036, 0.0057])

    TFÚ	vif_scoreÚtotalÚ
sigma_n_sqÚ	reduction©ÚmeanÚnoneÚkwargsÚreturnNc                 óÞ   •— t        ‰| �  di |¤Ž t        |t        t        f«      r|dk  rt        d|› �«      ‚|dvrt        d|› �«      ‚|| _        || _        | j                  dg d ¬«       y )Nr   zIArgument `sigma_n_sq` is expected to be a positive float or int, but got r   z7Argument `reduction` must be 'mean' or 'none', but got r   )ÚdefaultÚdist_reduce_fx© )	ÚsuperÚ__init__Ú
isinstanceÚfloatÚintÚ
ValueErrorr   r   Ú	add_state)Úselfr   r   r   Ú	__class__s       €úk/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchmetrics/image/vif.pyr   z"VisualInformationFidelity.__init__H   sz   ø€ Ü‰ÑÑ"˜6Ò"ä˜*¤u¬c lÔ3°zÀA²~ÜÐhÐisÐhtÐuÓvÐvàÐ,Ñ,ÜÐVÐW`ÐVaÐbÓcÐcà$ˆŒØ"ˆŒØ�‰�{¨B¸tˆÕDó    ÚpredsÚtargetc                 óz  — |j                  d«      }t        |«      D �cg c]3  }t        |dd…|dd…dd…f   |dd…|dd…dd…f   | j                  «      ‘Œ5 }}|dkD  r)t	        j
                  t	        j                  |«      d«      nt	        j                  |«      }| j                  j                  |«       yc c}w )z*Update state with predictions and targets.é   Nr   )
ÚsizeÚranger   r   Útorchr   ÚstackÚcatr   Úappend)r    r$   r%   ÚchannelsÚiÚvif_per_channels         r"   Úupdatez VisualInformationFidelity.updateU   s¡   € à—:‘:˜a“=ˆä^cÐdlÓ^mö
ØYZÔ˜U¢1 aªªA :Ñ.°²q¸!ºQÂ°zÑ0BÀDÇOÁOÕTð
ˆð 
ð JRÐTUÊœ%Ÿ*™*¤U§[¡[°Ó%AÀ1ÔEÔ[`×[dÑ[dÐetÓ[uˆØ�‰×Ñ˜oÕ.ùò	
s   Ÿ8B8c                 ón   — t        | j                  «      }| j                  dk(  r|j                  «       S |S )zCompute VIF over state.r   )r	   r   r   r   )r    r   s     r"   Úcomputez!VisualInformationFidelity.compute^   s/   € ä  §¡Ó0ˆ	Ø�>‰>˜VÒ#Ø—>‘>Ó#Ð#ØÐr#   )g       @r   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__Úis_differentiableÚhigher_is_betterÚfull_state_updater   r   Ú__annotations__r   r   r   r   r1   r3   Ú__classcell__)r!   s   @r"   r   r      s   ø… ñ%ðN ÐØÐØÐà�F‰|ÓØƒMñE 5ð E¸7À>Ñ;Rð EÐgjð EÐosõ Eð/˜Fð /¨Fð /°tó /ð˜÷ r#   r   )Útypingr   r   r*   r   Útyping_extensionsr   Ú!torchmetrics.functional.image.vifr   Útorchmetrics.metricr   Útorchmetrics.utilities.datar	   r   r   r#   r"   ú<module>rB      s*   ð÷ ã Ý Ý %å >Ý &Ý 4ôJ õ Jr#   