Ë
    þÍ:j$  ã                   óœ   — d dl mZ d dlmZmZmZ d dlmZmZ d dl	m
Z
 d dlmZmZ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ÚOptionalÚUnion)ÚTensorÚtensor)ÚLiteral)ÚALLOWED_MULTIOUTPUTÚ_explained_variance_computeÚ_explained_variance_update)ÚMetric)Ú_MATPLOTLIB_AVAILABLE)Ú_AX_TYPEÚ_PLOT_OUT_TYPEzExplainedVariance.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d<   eed<   eed<   eed<   eed<   	 dded   deddfˆ fd„Zdededdfd„Zde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 )ÚExplainedVarianceaˆ  Compute `explained variance`_.

    .. math:: \text{ExplainedVariance} = 1 - \frac{\text{Var}(y - \hat{y})}{\text{Var}(y)}

    Where :math:`y` is a tensor of target values, and :math:`\hat{y}` is a tensor of predictions.

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

    - ``preds`` (:class:`~torch.Tensor`): Predictions from model in float tensor
      with shape ``(N,)`` or ``(N, ...)`` (multioutput)
    - ``target`` (:class:`~torch.Tensor`): Ground truth values in long tensor
      with shape ``(N,)`` or ``(N, ...)`` (multioutput)

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

    - ``explained_variance`` (:class:`~torch.Tensor`): A tensor with the explained variance(s)

    In the case of multioutput, as default the variances will be uniformly averaged over the additional dimensions.
    Please see argument ``multioutput`` for changing this behavior.

    Args:
        multioutput:
            Defines aggregation in the case of multiple output scores. Can be one
            of the following strings (default is ``'uniform_average'``.):

            * ``'raw_values'`` returns full set of scores
            * ``'uniform_average'`` scores are uniformly averaged
            * ``'variance_weighted'`` scores are weighted by their individual variances

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

    Raises:
        ValueError:
            If ``multioutput`` is not one of ``"raw_values"``, ``"uniform_average"`` or ``"variance_weighted"``.

    Example:
        >>> from torch import tensor
        >>> from torchmetrics.regression import ExplainedVariance
        >>> target = tensor([3, -0.5, 2, 7])
        >>> preds = tensor([2.5, 0.0, 2, 8])
        >>> explained_variance = ExplainedVariance()
        >>> explained_variance(preds, target)
        tensor(0.9572)

        >>> target = tensor([[0.5, 1], [-1, 1], [7, -6]])
        >>> preds = tensor([[0, 2], [-1, 2], [8, -5]])
        >>> explained_variance = ExplainedVariance(multioutput='raw_values')
        >>> explained_variance(preds, target)
        tensor([0.9677, 1.0000])

    TÚis_differentiableÚhigher_is_betterFÚfull_state_updateç        Úplot_lower_boundg      ð?Úplot_upper_boundÚnum_obsÚ	sum_errorÚsum_squared_errorÚ
sum_targetÚsum_squared_targetÚmultioutput)Ú
raw_valuesÚuniform_averageÚvariance_weightedÚkwargsÚreturnNc                 óˆ  •— t        ‰| �  d
i |¤Ž |t        vrt        dt        › �«      ‚|| _        | j                  dt        d«      d¬«       | j                  dt        d«      d¬«       | j                  dt        d«      d¬«       | j                  dt        d«      d¬«       | j                  d	t        d«      d¬«       y )NzFInvalid input to argument `multioutput`. Choose one of the following: r   r   Úsum)ÚdefaultÚdist_reduce_fxr   r   r   r   © )ÚsuperÚ__init__r
   Ú
ValueErrorr   Ú	add_stater   )Úselfr   r"   Ú	__class__s      €ú/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchmetrics/regression/explained_variance.pyr*   zExplainedVariance.__init__b   s³   ø€ ô
 	‰ÑÑ"˜6Ò"àÔ1Ñ1ÜØXÔYlÐXmÐnóð ð 'ˆÔØ�‰�{¬F°3«KÈˆÔNØ�‰Ð*´F¸3³KÐPUˆÔVØ�‰�|¬V°C«[ÈˆÔOØ�‰Ð+´V¸C³[ÐQVˆÔWØ�‰�y¬&°«+ÀeˆÕLó    ÚpredsÚtargetc                 óð   — t        ||«      \  }}}}}| j                  |z   | _        | j                  |z   | _        | j                  |z   | _        | j                  |z   | _        | j
                  |z   | _        y)z*Update state with predictions and targets.N)r   r   r   r   r   r   )r-   r1   r2   r   r   r   r   r   s           r/   ÚupdatezExplainedVariance.updatet   sy   € äPjØ�6óQ
ÑMˆ�Ð-¨zÐ;Mð —|‘| gÑ-ˆŒØŸ™¨)Ñ3ˆŒØ!%×!7Ñ!7Ð:KÑ!KˆÔØŸ/™/¨JÑ6ˆŒØ"&×"9Ñ"9Ð<NÑ"NˆÕr0   c                 óš   — t        | j                  | j                  | j                  | j                  | j
                  | j                  «      S )z&Compute explained variance over state.)r   r   r   r   r   r   r   )r-   s    r/   ÚcomputezExplainedVariance.compute   s?   € ä*Ø�L‰LØ�N‰NØ×"Ñ"Ø�O‰OØ×#Ñ#Ø×Ñó
ð 	
r0   Ú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

            >>> from torch import randn
            >>> # Example plotting a single value
            >>> from torchmetrics.regression import ExplainedVariance
            >>> metric = ExplainedVariance()
            >>> metric.update(randn(10,), randn(10,))
            >>> fig_, ax_ = metric.plot()

        .. plot::
            :scale: 75

            >>> from torch import randn
            >>> # Example plotting multiple values
            >>> from torchmetrics.regression import ExplainedVariance
            >>> metric = ExplainedVariance()
            >>> values = []
            >>> for _ in range(10):
            ...     values.append(metric(randn(10,), randn(10,)))
            >>> fig, ax = metric.plot(values)

        )Ú_plot)r-   r7   r8   s      r/   ÚplotzExplainedVariance.plotŠ   s   € ðP �z‰z˜#˜rÓ"Ð"r0   )r    )NN)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   ÚboolÚ__annotations__r   r   r   Úfloatr   r   r	   r   r*   r4   r   r   r6   r   r   r   r;   Ú__classcell__)r.   s   @r/   r   r   !   s  ø… ñ2ðh #Ð�tÓ"Ø!Ð�dÓ!Ø#Ð�tÓ#Ø!Ð�eÓ!Ø!Ð�eÓ!àƒOØÓØÓØÓØÓð VgñMàÐQÑRðMð ðMð 
õ	Mð$	O˜Fð 	O¨Fð 	O°tó 	Oð	
˜˜v x°Ñ'7Ð7Ñ8ó 	
ð _cñ(#Ø˜E &¨(°6Ñ*:Ð":Ñ;Ñ<ð(#ØIQÐRZÑI[ð(#à	÷(#r0   r   N)Úcollections.abcr   Útypingr   r   r   Útorchr   r   Útyping_extensionsr	   Ú5torchmetrics.functional.regression.explained_variancer
   r   r   Útorchmetrics.metricr   Útorchmetrics.utilities.importsr   Útorchmetrics.utilities.plotr   r   Ú__doctest_skip__r   r(   r0   r/   ú<module>rM      sE   ðõ %ß 'Ñ 'ç  Ý %÷ñ õ
 'Ý @ß @áØ0Ð1ÐôQ#˜õ Q#r0   