Ë
    þÍ: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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mZ esdgZerdd„Zer ee«      s	ddgZnddgZ G d„ de«      Zy)é    )ÚSequence)ÚAnyÚClassVarÚOptionalÚUnion)ÚTensor)ÚLiteral)Ú_LPIPSÚ_lpips_computeÚ_lpips_updateÚ_NoTrainLpips)ÚMetric)Údim_zero_cat)Ú_SKIP_SLOW_DOCTESTÚ_try_proceed_with_timeout)Ú_MATPLOTLIB_AVAILABLEÚ_TORCHVISION_AVAILABLE)Ú_AX_TYPEÚ_PLOT_OUT_TYPEz*LearnedPerceptualImagePatchSimilarity.plotNc                  ó   — t        dd¬«       y )NTÚvgg)Ú
pretrainedÚnet)r
   © ó    úl/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchmetrics/image/lpip.pyÚ_download_lpipsr       s   € Ü˜$ EÖ*r   Ú%LearnedPerceptualImagePatchSimilarityc                   ó(  ‡ — 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<   dZeed<   dgZeee      ed<   	 	 	 dded   deed      de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 )!r   a–  The Learned Perceptual Image Patch Similarity (`LPIPS_`) calculates perceptual similarity between two images.

    LPIPS essentially computes the similarity between the activations of two image patches for some pre-defined network.
    This measure has been shown to match human perception well. A low LPIPS score means that image patches are
    perceptual similar.

    Both input image patches are expected to have shape ``(N, 3, H, W)``. The minimum size of `H, W` depends on the
    chosen backbone (see `net_type` arg).

    .. hint::
        Using this metrics requires you to have ``torchvision`` package installed. Either install as
        ``pip install torchmetrics[image]`` or ``pip install torchvision``.

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

    - ``img1`` (:class:`~torch.Tensor`): tensor with images of shape ``(N, 3, H, W)``
    - ``img2`` (:class:`~torch.Tensor`): tensor with images of shape ``(N, 3, H, W)``

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

    - ``lpips`` (:class:`~torch.Tensor`): returns float scalar tensor with average LPIPS value over samples

    Args:
        net_type: str indicating backbone network type to use. Choose between `'alex'`, `'vgg'` or `'squeeze'`
        reduction: str indicating how to reduce over the batch dimension. Choose between `'sum'`, `'mean'`,`'none'`
            or `None`.
        normalize: by default this is ``False`` meaning that the input is expected to be in the [-1,1] range. If set
            to ``True`` will instead expect input to be in the ``[0,1]`` range.
        kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info.

    Raises:
        ModuleNotFoundError:
            If ``torchvision`` package is not installed
        ValueError:
            If ``net_type`` is not one of ``"vgg"``, ``"alex"`` or ``"squeeze"``
        ValueError:
            If ``reduction`` is not one of ``"mean"`` or ``"sum"``

    Example:
        >>> from torch import rand
        >>> from torchmetrics.image.lpip import LearnedPerceptualImagePatchSimilarity
        >>> lpips = LearnedPerceptualImagePatchSimilarity(net_type='squeeze')
        >>> # LPIPS needs the images to be in the [-1, 1] range.
        >>> img1 = (rand(10, 3, 100, 100) * 2) - 1
        >>> img2 = (rand(10, 3, 100, 100) * 2) - 1
        >>> lpips(img1, img2)
        tensor(0.1024)

        >>> from torch import rand, Generator
        >>> from torchmetrics.image.lpip import LearnedPerceptualImagePatchSimilarity
        >>> gen = Generator().manual_seed(42)
        >>> lpips = LearnedPerceptualImagePatchSimilarity(net_type='squeeze', reduction='none')
        >>> # LPIPS needs the images to be in the [-1, 1] range.
        >>> img1 = (rand(2, 3, 100, 100, generator=gen) * 2) - 1
        >>> img2 = (rand(2, 3, 100, 100, generator=gen) * 2) - 1
        >>> lpips(img1, img2)
        tensor([0.1024, 0.0938])

    TÚis_differentiableFÚhigher_is_betterÚfull_state_updateg        Úplot_lower_boundg      ð?Úplot_upper_boundÚ
all_scoresr   Úfeature_networkÚ__jit_ignored_attributes__Únet_type©r   ÚalexÚsqueezeÚ	reduction)ÚsumÚmeanÚnoneÚ	normalizeÚkwargsÚreturnNc                 óF  •— t        ‰| �  di |¤Ž t        st        d«      ‚d}||vrt	        d|› d|› d�«      ‚t        |¬«      | _        d}||vrt	        d|› d|› �«      ‚|| _        t        |t        «      st	        d	|› �«      ‚|| _
        | j                  d
g d ¬«       y )Nz†LPIPS metric requires that torchvision is installed. Either install as `pip install torchmetrics[image]` or `pip install torchvision`.r)   z#Argument `net_type` must be one of z
, but got ú.)r   )r.   r-   r/   Nz$Argument `reduction` must be one of z/Argument `normalize` should be an bool but got r%   )ÚdefaultÚdist_reduce_fxr   )ÚsuperÚ__init__r   ÚModuleNotFoundErrorÚ
ValueErrorr   r   r,   Ú
isinstanceÚboolr0   Ú	add_state)Úselfr(   r,   r0   r1   Úvalid_net_typeÚvalid_reductionÚ	__class__s          €r   r8   z.LearnedPerceptualImagePatchSimilarity.__init__r   sÍ   ø€ ô 	‰ÑÑ"˜6Ò"å%Ü%ðeóð ð
 4ˆØ˜>Ñ)ÜÐBÀ>ÐBRÐR\Ð]eÐ\fÐfgÐhÓiÐiÜ  XÔ.ˆŒà7ˆØ˜OÑ+ÜÐCÀOÐCTÐT^Ð_hÐ^iÐjÓkÐkØ"ˆŒä˜)¤TÔ*ÜÐNÈyÈkÐZÓ[Ð[Ø"ˆŒà�‰�|¨RÀˆÕEr   Úimg1Úimg2c                 ó€   — t        ||| j                  | j                  ¬«      }| j                  j	                  |«       y)z(Update internal states with lpips score.)r   r0   N)r   r   r0   r%   Úappend)r>   rB   rC   Úlosss       r   Úupdatez,LearnedPerceptualImagePatchSimilarity.update‘   s,   € ä˜T 4¨T¯X©XÀÇÁÔPˆØ�‰×Ñ˜tÕ$r   c                 óZ   — t        | j                  «      }t        || j                  ¬«      S )z+Compute final perceptual similarity metric.)r,   )r   r%   r   r,   )r>   Úscoress     r   Úcomputez-LearnedPerceptualImagePatchSimilarity.compute–   s!   € ä˜dŸo™oÓ.ˆÜ˜f°·±Ô?Ð?r   Ú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.lpip import LearnedPerceptualImagePatchSimilarity
            >>> metric = LearnedPerceptualImagePatchSimilarity(net_type='squeeze')
            >>> metric.update(torch.rand(10, 3, 100, 100), torch.rand(10, 3, 100, 100))
            >>> fig_, ax_ = metric.plot()

        .. plot::
            :scale: 75

            >>> # Example plotting multiple values
            >>> import torch
            >>> from torchmetrics.image.lpip import LearnedPerceptualImagePatchSimilarity
            >>> metric = LearnedPerceptualImagePatchSimilarity(net_type='squeeze')
            >>> values = [ ]
            >>> for _ in range(3):
            ...     values.append(metric(torch.rand(10, 3, 100, 100), torch.rand(10, 3, 100, 100)))
            >>> fig_, ax_ = metric.plot(values)

        )Ú_plot)r>   rK   rL   s      r   Úplotz*LearnedPerceptualImagePatchSimilarity.plot›   s   € ðP �z‰z˜#˜rÓ"Ð"r   )r*   r.   F)NN)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r    r<   Ú__annotations__r!   r"   r#   Úfloatr$   Úlistr   r&   Ústrr'   r   r	   r   r   r8   rG   rJ   r   r   r   r   rO   Ú__classcell__)rA   s   @r   r   r   )   s+  ø… ñ:ðx #Ð�tÓ"Ø"Ð�dÓ"Ø#Ð�tÓ#Ø!Ð�eÓ!Ø!Ð�eÓ!à�V‘ÓØ €O�SÓ ð 8=°gÐ ¨¨c©Ñ 3Ó=ð 7=Ø>DØñ	FàÐ2Ñ3ðFð ˜GÐ$9Ñ:Ñ;ðFð ð	Fð
 ðFð 
õFð>%˜6ð %¨ð %°Dó %ð
@˜ó @ð _cñ(#Ø˜E &¨(°6Ñ*:Ð":Ñ;Ñ<ð(#ØIQÐRZÑI[ð(#à	÷(#r   )r2   N) Úcollections.abcr   Útypingr   r   r   r   Útorchr   Útyping_extensionsr	   Ú#torchmetrics.functional.image.lpipsr
   r   r   r   Útorchmetrics.metricr   Útorchmetrics.utilitiesr   Útorchmetrics.utilities.checksr   r   Útorchmetrics.utilities.importsr   r   Útorchmetrics.utilities.plotr   r   Ú__doctest_skip__r   r   r   r   r   ú<module>rd      sr   ðõ %ß 1Ó 1å Ý %ç dÓ dÝ &Ý /ß Wß Xß @áØDÐEÐáó+ñ Ñ";¸OÔ"LØCÐEqÐrÑà?ÐAmÐnÐôZ#¨Fõ Z#r   