Ë
    þÍ:j)&  ã                   ó  — d dl mZ d dlmZmZmZmZ d dlZd dl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sd	gZ	 dd
ej                  dej                  dej                  dej                  dej                  dej                  dej                  dej                  dedeej                  ej                  ej                  ej                  ej                  ej                  ej                  ej                  f   fd„Z G d„ de«      Zy)é    )ÚSequence)ÚAnyÚListÚOptionalÚUnionN)ÚTensor)Ú_pearson_corrcoef_computeÚ_pearson_corrcoef_update)ÚMetric)Ú_MATPLOTLIB_AVAILABLE)Ú_AX_TYPEÚ_PLOT_OUT_TYPEzPearsonCorrCoef.plotÚmeans_xÚmeans_yÚ
maxs_abs_xÚ
maxs_abs_yÚvars_xÚvars_yÚcorrs_xyÚnbsÚepsÚreturnc	           	      óÎ  — t        | «      dk(  r"| d   |d   |d   |d   |d   |d   |d   |d   fS | d   }	|d   }
|d   }|d   }|d   }|d   }|d   }|d   }t        dt        | «      «      D ]ì  }| |   }||   }||   }||   }||   }||   }||   }||   }t        j                  t        j                  ||«      ||z   |«      }||	z  ||z  z   |z  }||
z  ||z  z   |z  }||z  |z  }||	z
  }||
z
  }||z   ||dz  z  z   } ||z   ||dz  z  z   }!||z   ||z  |z  z   }"t        j
                  ||«      }#t        j
                  ||«      }$|}	|}
|#}|$}| }|!}|"}|}Œî #$ !"fS )a¥  Aggregate the statistics from multiple devices.

    Formula taken from here: `Parallel algorithm for calculating variance
    <https://en.wikipedia.org/wiki/Algorithms_for_calculating_variance#Parallel_algorithm>`_

    We use `eps` to avoid division by zero when `n1` and `n2` are both zero. Generally, the value of `eps` should not
    matter, as if `n1` and `n2` are both zero, all the states will also be zero.

    é   r   é   )ÚlenÚrangeÚtorchÚwhereÚ
logical_orÚmaximum)%r   r   r   r   r   r   r   r   r   Úmx1Úmy1Úmax1Úmay1Úvx1Úvy1Úcxy1Ún1ÚiÚmx2Úmy2Úmax2Úmay2Úvx2Úvy2Úcxy2Ún2ÚnbÚmean_xÚmean_yÚn12_bÚdelta_xÚdelta_yÚvar_xÚvar_yÚcorr_xyÚmax_abs_dev_xÚmax_abs_dev_ys%                                        út/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchmetrics/regression/pearson.pyÚ_final_aggregationr?      s  € ô, ˆ7ƒ|�qÒØ�q‰z˜7 1™: z°!¡}°jÀ±mÀVÈAÁYÐPVÐWXÑPYÐ[cÐdeÑ[fÐhkÐlmÑhnÐnÐnØ
�!‰*€CØ
�!‰*€CØ�a‰=€DØ�a‰=€DØ
�‰)€CØ
�‰)€CØ�A‰;€DØ	ˆQ‰€BÜ�1”c˜'“lÓ#ò #ˆØ�a‰jˆØ�a‰jˆØ˜!‰}ˆØ˜!‰}ˆØ�Q‰iˆØ�Q‰iˆØ˜‰{ˆØ�‰Vˆä�[‰[œ×)Ñ)¨"¨bÓ1°2¸±7¸CÓ@ˆà�s‘(˜R #™XÑ%¨Ñ+ˆà�s‘(˜R #™XÑ%¨Ñ+ˆà�R‘˜"‘ˆØ˜‘)ˆØ˜‘)ˆà�c‘	˜E G¨Q¡JÑ.Ñ.ˆà�c‘	˜E G¨Q¡JÑ.Ñ.ˆà˜‘+ ¨¡°'Ñ 9Ñ9ˆÜŸ™ d¨DÓ1ˆÜŸ™ d¨DÓ1ˆàˆØˆØˆØˆØˆØˆØˆØ‰ðG#ðH �6˜=¨-¸ÀÀwÐPRÐRÐRó    c                   óF  ‡ — e Zd ZU dZdZeed<   dZee   ed<   dZ	eed<   dZ
eed<   d	Zeed
<   ee   ed<   ee   ed<   eed<   eed<   eed<   eed<   eed<   eed<   eed<   eed<   	 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 ) ÚPearsonCorrCoefa  Compute `Pearson Correlation Coefficient`_.

    .. math::
        P_{corr}(x,y) = \frac{cov(x,y)}{\sigma_x \sigma_y}

    Where :math:`y` is a tensor of target values, and :math:`x` is a tensor of predictions.

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

    - ``preds`` (:class:`~torch.Tensor`): either single output float tensor with shape ``(N,)``
      or multioutput float tensor of shape ``(N,d)``
    - ``target`` (:class:`~torch.Tensor`): either single output tensor with shape ``(N,)``
      or multioutput tensor of shape ``(N,d)``

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

    - ``pearson`` (:class:`~torch.Tensor`): A tensor with the Pearson Correlation Coefficient

    Args:
        num_outputs: Number of outputs in multioutput setting
        kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info.

    Example (single output regression):
        >>> from torchmetrics.regression import PearsonCorrCoef
        >>> target = torch.tensor([3, -0.5, 2, 7])
        >>> preds = torch.tensor([2.5, 0.0, 2, 8])
        >>> pearson = PearsonCorrCoef()
        >>> pearson(preds, target)
        tensor(0.9849)

    Example (multi output regression):
        >>> from torchmetrics.regression import PearsonCorrCoef
        >>> target = torch.tensor([[3, -0.5], [2, 7]])
        >>> preds = torch.tensor([[2.5, 0.0], [2, 8]])
        >>> pearson = PearsonCorrCoef(num_outputs=2)
        >>> pearson(preds, target)
        tensor([1., 1.])

    TÚis_differentiableNÚhigher_is_betterÚfull_state_updateg      ð¿Úplot_lower_boundg      ð?Úplot_upper_boundÚpredsÚtargetr4   r5   r<   r=   r9   r:   r;   Ún_totalÚnum_outputsÚkwargsr   c                 ó‚  •— t        ‰| �  di |¤Ž t        |t        «      s|dk  rt	        d«      ‚|| _        | j                  dt        j                  | j
                  «      d ¬«       | j                  dt        j                  | j
                  «      d ¬«       | j                  dt        j                  | j
                  «      d ¬«       | j                  dt        j                  | j
                  «      d ¬«       | j                  dt        j                  | j
                  «      d ¬«       | j                  d	t        j                  | j
                  «      d ¬«       | j                  d
t        j                  | j
                  «      d ¬«       | j                  dt        j                  | j
                  «      d ¬«       y )Nr   zQExpected argument `num_outputs` to be an int larger than 0, but got {num_outputs}r4   )ÚdefaultÚdist_reduce_fxr5   r<   r=   r9   r:   r;   rJ   © )	ÚsuperÚ__init__Ú
isinstanceÚintÚ
ValueErrorrK   Ú	add_stater   Úzeros)ÚselfrK   rL   Ú	__class__s      €r>   rR   zPearsonCorrCoef.__init__�   sI  ø€ ô
 	‰ÑÑ"˜6Ò"Ü˜+¤sÔ+°¸a²ÜÐpÓqÐqØ&ˆÔà�‰�x¬¯©°T×5EÑ5EÓ)FÐW[ˆÔ\Ø�‰�x¬¯©°T×5EÑ5EÓ)FÐW[ˆÔ\Ø�‰�´·±¸D×<LÑ<LÓ0MÐ^bˆÔcØ�‰�´·±¸D×<LÑ<LÓ0MÐ^bˆÔcØ�‰�w¬¯©°D×4DÑ4DÓ(EÐVZˆÔ[Ø�‰�w¬¯©°D×4DÑ4DÓ(EÐVZˆÔ[Ø�‰�y¬%¯+©+°d×6FÑ6FÓ*GÐX\ˆÔ]Ø�‰�y¬%¯+©+°d×6FÑ6FÓ*GÐX\ˆÕ]r@   c                 óF  — t        ||| j                  | j                  | j                  | j                  | j
                  | j                  | j                  | j                  | j                  ¬«      \  | _        | _        | _        | _        | _        | _        | _        | _        y)z*Update state with predictions and targets.)rH   rI   r4   r5   r<   r=   r9   r:   r;   Ú	num_priorrK   N)
r
   r4   r5   r<   r=   r9   r:   r;   rJ   rK   )rX   rH   rI   s      r>   ÚupdatezPearsonCorrCoef.update°   s†   € ô %ØØØ—;‘;Ø—;‘;Ø×,Ñ,Ø×,Ñ,Ø—*‘*Ø—*‘*Ø—L‘LØ—l‘lØ×(Ñ(ô
ñ		
ØŒKØŒKØÔØÔØŒJØŒJØŒLØ�Lr@   c           
      ó4  — | j                   dk(  r| j                  j                  «       dkD  s(| j                   dkD  r†| j                  j                  dkD  rmt	        | j                  | j
                  | j                  | j                  | j                  | j                  | j                  | j                  ¬«      \  }}}}}}}}nH| j                  }| j                  }| j                  }| j                  }| j                  }| j                  }t        ||||||«      S )z3Compute pearson correlation coefficient over state.r   )r   r   r   r   r   r   r   r   )rK   r4   ÚnumelÚndimr?   r5   r<   r=   r9   r:   r;   rJ   r	   )rX   Ú_r<   r=   r9   r:   r;   rJ   s           r>   ÚcomputezPearsonCorrCoef.computeÉ   sò   € à×Ñ Ò! d§k¡k×&7Ñ&7Ó&9¸AÒ&=À4×CSÑCSÐVWÒCWÐ\`×\gÑ\g×\lÑ\lÐopÒ\päQcØŸ™ØŸ™Ø×-Ñ-Ø×-Ñ-Ø—z‘zØ—z‘zØŸ™Ø—L‘Lô	RÑNˆAˆq�- °°u¸gÁwð !×.Ñ.ˆMØ ×.Ñ.ˆMØ—J‘JˆEØ—J‘JˆEØ—l‘lˆGØ—l‘lˆGÜ(¨¸ÀuÈeÐU\Ð^eÓfÐfr@   Ú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 PearsonCorrCoef
            >>> metric = PearsonCorrCoef()
            >>> 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 PearsonCorrCoef
            >>> metric = PearsonCorrCoef()
            >>> values = []
            >>> for _ in range(10):
            ...     values.append(metric(randn(10,), randn(10,)))
            >>> fig, ax = metric.plot(values)

        )Ú_plot)rX   rb   rc   s      r>   ÚplotzPearsonCorrCoef.plotà   s   € ðP �z‰z˜#˜rÓ"Ð"r@   )r   )NN)Ú__name__Ú
__module__Ú__qualname__Ú__doc__rC   ÚboolÚ__annotations__rD   r   rE   rF   ÚfloatrG   r   r   rT   r   rR   r\   ra   r   r   r   r   rf   Ú__classcell__)rY   s   @r>   rB   rB   d   s  ø… ñ&ðP #Ð�tÓ"Ø'+Ð�h˜t‘nÓ+Ø"Ð�tÓ"Ø"Ð�eÓ"Ø!Ð�eÓ!Ø�‰<ÓØ�‰LÓØƒNØƒNØÓØÓØƒMØƒMØƒOØƒOð ñ^àð^ð ð^ð 
õ	^ð&
˜Fð 
¨Fð 
°tó 
ð2g˜ó gð0 _cñ(#Ø˜E &¨(°6Ñ*:Ð":Ñ;Ñ<ð(#ØIQÐRZÑI[ð(#à	÷(#r@   rB   )g»½×Ùß|Û=)Úcollections.abcr   Útypingr   r   r   r   r   r   Ú*torchmetrics.functional.regression.pearsonr	   r
   Útorchmetrics.metricr   Útorchmetrics.utilities.importsr   Útorchmetrics.utilities.plotr   r   Ú__doctest_skip__rm   Útupler?   rB   rP   r@   r>   ú<module>rw      s  ðõ %ß -Ó -ã Ý ç jÝ &Ý @ß @áØ.Ð/Ðð ñDSØ�\‰\ðDSà�\‰\ðDSð —‘ðDSð —‘ð	DSð
 �L‰LðDSð �L‰LðDSð �l‰lðDSð 
�‰ðDSð 
ðDSð Ø	‡L�L�%—,‘, §¡¨e¯l©l¸E¿L¹LÈ%Ï,É,ÐX]×XdÑXdÐfk×frÑfrÐrñóDSôNd#�fõ d#r@   