a
    d                     @   s   d dl mZmZmZmZ d dl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mZ esrdgZG d	d
 d
eZdS )    )AnyOptionalSequenceUnionN)Tensortensor)_r2_score_compute_r2_score_update)Metric)_MATPLOTLIB_AVAILABLE)_AX_TYPE_PLOT_OUT_TYPER2Score.plotc                       s   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< deeeedd fddZeeddddZedddZd eeeee f  ee edddZ  ZS )!R2Scorea  Compute r2 score also known as `R2 Score_Coefficient Determination`_.

    .. math:: R^2 = 1 - \frac{SS_{res}}{SS_{tot}}

    where :math:`SS_{res}=\sum_i (y_i - f(x_i))^2` is the sum of residual squares, and
    :math:`SS_{tot}=\sum_i (y_i - \bar{y})^2` is total sum of squares. Can also calculate
    adjusted r2 score given by

    .. math:: R^2_{adj} = 1 - \frac{(1-R^2)(n-1)}{n-k-1}

    where the parameter :math:`k` (the number of independent regressors) should be provided as the `adjusted` argument.
    The score is only proper defined when :math:`SS_{tot}\neq 0`, which can happen for near constant targets. In this
    case a score of 0 is returned. By definition the score is bounded between 0 and 1, where 1 corresponds to the
    predictions exactly matching the targets.

    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, M)`` (multioutput)
    - ``target`` (:class:`~torch.Tensor`): Ground truth values in float tensor with shape ``(N,)``
      or ``(N, M)`` (multioutput)

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

    - ``r2score`` (:class:`~torch.Tensor`): A tensor with the r2 score(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:
        num_outputs: Number of outputs in multioutput setting
        adjusted: number of independent regressors for calculating adjusted r2 score.
        multioutput: Defines aggregation in the case of multiple output scores. Can be one of the following strings:

            * ``'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 ``adjusted`` parameter is not an integer larger or equal to 0.
        ValueError:
            If ``multioutput`` is not one of ``"raw_values"``, ``"uniform_average"`` or ``"variance_weighted"``.

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

        >>> target = torch.tensor([[0.5, 1], [-1, 1], [7, -6]])
        >>> preds = torch.tensor([[0, 2], [-1, 2], [8, -5]])
        >>> r2score = R2Score(num_outputs=2, multioutput='raw_values')
        >>> r2score(preds, target)
        tensor([0.9654, 0.9082])

    Tis_differentiablehigher_is_betterFfull_state_updateg        plot_lower_boundg      ?plot_upper_boundsum_squared_error	sum_errorresidualtotal   r   uniform_averageN)num_outputsadjustedmultioutputkwargsreturnc                    s   t  jf i | || _|dk s*t|ts2td|| _d}||vrRtd| || _| jdt	
| jdd | jdt	
| jdd | jd	t	
| jdd | jd
tddd d S )Nr   z?`adjusted` parameter should be an integer larger or equal to 0.)Z
raw_valuesr   Zvariance_weightedzFInvalid input to argument `multioutput`. Choose one of the following: r   sum)defaultZdist_reduce_fxr   r   r   )super__init__r   
isinstanceint
ValueErrorr   r   Z	add_statetorchzerosr   )selfr   r   r   r   Zallowed_multioutput	__class__ c/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/torchmetrics/regression/r2.pyr#   d   s    zR2Score.__init__)predstargetr   c                 C   sN   t ||\}}}}|  j|7  _|  j|7  _|  j|7  _|  j|7  _dS )z*Update state with predictions and targets.N)r	   r   r   r   r   )r)   r.   r/   r   r   r   r   r,   r,   r-   update   s
    zR2Score.update)r   c                 C   s   t | j| j| j| j| j| jS )z(Compute r2 score over the metric states.)r   r   r   r   r   r   r   )r)   r,   r,   r-   compute   s    zR2Score.compute)valaxr   c                 C   s   |  ||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 R2Score
            >>> metric = R2Score()
            >>> 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 R2Score
            >>> metric = R2Score()
            >>> values = []
            >>> for _ in range(10):
            ...     values.append(metric(randn(10,), randn(10,)))
            >>> fig, ax = metric.plot(values)

        )Z_plot)r)   r2   r3   r,   r,   r-   plot   s    (r   )r   r   r   )NN)__name__
__module____qualname____doc__r   bool__annotations__r   r   r   floatr   r   r%   strr   r#   r0   r1   r   r   r   r   r   r4   __classcell__r,   r,   r*   r-   r      s6   
<   	 r   )typingr   r   r   r   r'   r   r   Z%torchmetrics.functional.regression.r2r   r	   Ztorchmetrics.metricr
   Ztorchmetrics.utilities.importsr   Ztorchmetrics.utilities.plotr   r   Z__doctest_skip__r   r,   r,   r,   r-   <module>   s   