Ë
    þÍ:j$  ã                   ó¾   — d dl mZ d dlmZ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 d d	lmZ d d
lmZmZ d dlmZmZ esdgZddgiZ G d„ de«      Zy)é    )ÚSequence)ÚAnyÚOptionalÚUnionN)ÚTensor)ÚModule)ÚNoTrainInceptionV3)ÚMetric)Úrank_zero_warn)Údim_zero_cat)Ú_MATPLOTLIB_AVAILABLEÚ_TORCH_FIDELITY_AVAILABLE)Ú_AX_TYPEÚ_PLOT_OUT_TYPEúInceptionScore.plot)ÚInceptionScorer   Útorch_fidelityc                   ó  ‡ — 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<   eed	<   eed
<   d
Zeed<   	 	 	 ddeeeef   dedededdf
ˆ fd„Zdeddfd„Zdeeef   fd„Z	 ddeeeee   f      dee   defd„Zˆ xZS )r   aÜ  Calculate the Inception Score (IS) which is used to access how realistic generated images are.

    .. math::
        IS = exp(\mathbb{E}_x KL(p(y | x ) || p(y)))

    where :math:`KL(p(y | x) || p(y))` is the KL divergence between the conditional distribution :math:`p(y|x)`
    and the marginal distribution :math:`p(y)`. Both the conditional and marginal distribution is calculated
    from features extracted from the images. The score is calculated on random splits of the images such that
    both a mean and standard deviation of the score are returned. The metric was originally proposed in
    `inception ref1`_.

    Using the default feature extraction (Inception v3 using the original weights from `inception ref2`_), the input
    is expected to be mini-batches of 3-channel RGB images of shape ``(3xHxW)``. If argument ``normalize``
    is ``True`` images are expected to be dtype ``float`` and have values in the ``[0,1]`` range, else if
    ``normalize`` is set to ``False`` images are expected to have dtype uint8 and take values in the ``[0, 255]``
    range. All images will be resized to 299 x 299 which is the size of the original training data.

    .. hint::
        Using this metric with the default feature extractor requires that ``torch-fidelity``
        is installed. Either install as ``pip install torchmetrics[image]`` or
        ``pip install torch-fidelity``

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

    - ``imgs`` (:class:`~torch.Tensor`): tensor with images feed to the feature extractor

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

    - ``inception_mean`` (:class:`~torch.Tensor`): float scalar tensor with mean inception score over subsets
    - ``inception_std`` (:class:`~torch.Tensor`): float scalar tensor with standard deviation of inception score
      over subsets

    Args:
        feature:
            Either an str, integer or ``nn.Module``:

            - an str or integer will indicate the inceptionv3 feature layer to choose. Can be one of the following:
              'logits_unbiased', 64, 192, 768, 2048
            - an ``nn.Module`` for using a custom feature extractor. Expects that its forward method returns
              an ``(N,d)`` matrix where ``N`` is the batch size and ``d`` is the feature size.

        splits: integer determining how many splits the inception score calculation should be split among
        kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info.

    Raises:
        ValueError:
            If ``feature`` is set to an ``str`` or ``int`` and ``torch-fidelity`` is not installed
        ValueError:
            If ``feature`` is set to an ``str`` or ``int`` and not one of ``('logits_unbiased', 64, 192, 768, 2048)``
        TypeError:
            If ``feature`` is not an ``str``, ``int`` or ``torch.nn.Module``

    Example:
        >>> from torch import rand
        >>> from torchmetrics.image.inception import InceptionScore
        >>> inception = InceptionScore()
        >>> # generate some images
        >>> imgs = torch.randint(0, 255, (100, 3, 299, 299), dtype=torch.uint8)
        >>> inception.update(imgs)
        >>> inception.compute()
        (tensor(1.0549), tensor(0.0121))

    FÚis_differentiableTÚhigher_is_betterÚfull_state_updateg        Úplot_lower_boundÚfeaturesÚ	inceptionÚfeature_networkÚfeatureÚsplitsÚ	normalizeÚkwargsÚreturnNc                 ó¼  •— t        ‰| �  di |¤Ž t        dt        «       t	        |t
        t        f«      rFt        st        d«      ‚d}||vrt        d|› d|› d�«      ‚t        dt        |«      g¬«      | _        n#t	        |t        «      r|| _        nt        d	«      ‚t	        |t        «      st        d
«      ‚|| _        || _        | j#                  dg d ¬«       y )NzMetric `InceptionScore` will save all extracted features in buffer. For large datasets this may lead to large memory footprint.z—InceptionScore metric requires that `Torch-fidelity` is installed. Either install as `pip install torchmetrics[image]` or `pip install torch-fidelity`.)Úlogits_unbiasedé@   éÀ   i   i   z3Integer input to argument `feature` must be one of z
, but got ú.zinception-v3-compat)ÚnameÚfeatures_listz'Got unknown input to argument `feature`z*Argument `normalize` expected to be a boolr   )Údist_reduce_fx© )ÚsuperÚ__init__r   ÚUserWarningÚ
isinstanceÚstrÚintr   ÚModuleNotFoundErrorÚ
ValueErrorr	   r   r   Ú	TypeErrorÚboolr   r   Ú	add_state)Úselfr   r   r   r   Úvalid_int_inputÚ	__class__s         €úq/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchmetrics/image/inception.pyr+   zInceptionScore.__init__m   sî   ø€ ô 	‰ÑÑ"˜6Ò"äðKäô	
ô �g¤¤S˜zÔ*Ý,Ü)ðlóð ð FˆOØ˜oÑ-Ü ØIÈ/ÐIZÐZdÐelÐdmÐmnÐoóð ô 0Ð5JÔ[^Ð_fÓ[gÐZhÔiˆD�NÜ˜¤Ô(Ø$ˆD�NäÐEÓFÐFä˜)¤TÔ*ÜÐIÓJÐJØ"ˆŒàˆŒØ�‰�z 2°dˆÕ;ó    Úimgsc                 óž   — | j                   r|dz  j                  «       n|}| j                  |«      }| j                  j	                  |«       y)z)Update the state with extracted features.éÿ   N)r   Úbyter   r   Úappend)r5   r:   r   s      r8   ÚupdatezInceptionScore.update•   s<   € à&*§n¢n��s‘
× Ñ Ô"¸$ˆØ—>‘> $Ó'ˆØ�‰×Ñ˜XÕ&r9   c           	      óä  — t        | j                  «      }t        j                  |j                  d   «      }||   }|j                  d¬«      }|j                  d¬«      }|j                  | j                  d¬«      }|j                  | j                  d¬«      }|D �cg c]  }|j                  dd¬«      ‘Œ }}t        |||«      D ���cg c]  \  }}}|||j                  «       z
  z  ‘Œ }	}}}|	D �
cg c]0  }
|
j                  d¬«      j                  «       j                  «       ‘Œ2 }	}
t        j                  |	«      }|j                  «       |j                  «       fS c c}w c c}}}w c c}
w )zCompute metric.r   é   )ÚdimT)rB   Úkeepdim)r   r   ÚtorchÚrandpermÚshapeÚsoftmaxÚlog_softmaxÚchunkr   ÚmeanÚzipÚlogÚsumÚexpÚstackÚstd)r5   r   ÚidxÚprobÚlog_probÚpÚ	mean_probÚlog_pÚm_pÚkl_ÚkÚkls               r8   ÚcomputezInceptionScore.compute›   s@  € ä §¡Ó.ˆä�n‰n˜XŸ^™^¨AÑ.Ó/ˆØ˜C‘=ˆð ×Ñ AÐÓ&ˆØ×'Ñ'¨AÐ'Ó.ˆð �z‰z˜$Ÿ+™+¨1ˆzÓ-ˆØ—>‘> $§+¡+°1�>Ó5ˆð ;?Ö?°Q�Q—V‘V ¨4�VÕ0Ð?ˆ	Ð?Ü<?ÀÀhÐPYÓ<Z×[Ð[©=¨1¨e°Sˆq�E˜CŸG™G›IÑ%Ó&Ð[ˆÒ[Ø25Ö6¨Qˆq�u‰u˜ˆu‹|× Ñ Ó"×&Ñ&Õ(Ð6ˆÐ6Ü�[‰[˜Óˆð �w‰w‹y˜"Ÿ&™&›(Ð"Ð"ùò @ùÜ[ùÚ6s   ÂE!Ã!E&Ã55E-ÚvalÚaxc                 óT   — |xs | j                  «       d   }| 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.inception import InceptionScore
            >>> metric = InceptionScore()
            >>> metric.update(torch.randint(0, 255, (50, 3, 299, 299), dtype=torch.uint8))
            >>> fig_, ax_ = metric.plot()  # the returned plot only shows the mean value by default

        .. plot::
            :scale: 75

            >>> # Example plotting multiple values
            >>> import torch
            >>> from torchmetrics.image.inception import InceptionScore
            >>> metric = InceptionScore()
            >>> values = [ ]
            >>> for _ in range(3):
            ...     # we index by 0 such that only the mean value is plotted
            ...     values.append(metric(torch.randint(0, 255, (50, 3, 299, 299), dtype=torch.uint8))[0])
            >>> fig_, ax_ = metric.plot(values)

        r   )r[   Ú_plot)r5   r\   r]   s      r8   ÚplotzInceptionScore.plot³   s+   € ðR Ò&�T—\‘\“^ AÑ&ˆØ�z‰z˜#˜rÓ"Ð"r9   )r"   é
   F)NN)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   r3   Ú__annotations__r   r   r   ÚfloatÚlistr   r   r.   r   r/   r   r+   r   r?   Útupler[   r   r   r   r   r`   Ú__classcell__)r7   s   @r8   r   r   #   sÿ   ø… ñ>ð@ $Ð�tÓ#Ø!Ð�dÓ!Ø#Ð�tÓ#Ø!Ð�eÓ!àƒNØÓØ&€O�SÓ&ð ,=ØØñ	&<à�s˜C Ð'Ñ(ð&<ð ð&<ð ð	&<ð
 ð&<ð 
õ&<ðP'˜6ð ' dó 'ð#˜˜v v˜~Ñ.ó #ð2 _cñ*#Ø˜E &¨(°6Ñ*:Ð":Ñ;Ñ<ð*#ØIQÐRZÑI[ð*#à	÷*#r9   r   )Úcollections.abcr   Útypingr   r   r   rD   r   Útorch.nnr   Útorchmetrics.image.fidr	   Útorchmetrics.metricr
   Útorchmetrics.utilitiesr   Útorchmetrics.utilities.datar   Útorchmetrics.utilities.importsr   r   Útorchmetrics.utilities.plotr   r   Ú__doctest_skip__Ú__doctest_requires__r   r)   r9   r8   ú<module>rv      sW   ðõ %ß 'Ñ 'ã Ý Ý å 5Ý &Ý 1Ý 4ß [ß @áØ-Ð.Ðð BÐDTÐCUÐVÐ ôz#�Võ z#r9   