Ë
    þÍ:jàn  ã                   óP  — 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
 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 esg d¢Z G d„ de«      Z G d„ de«      Z G d„ de«      Z G d„ de«      Z G d„ de«      Z G d„ de«      Z G d„ de«      Z G d„ de«      Z y)é    )ÚSequence)ÚAnyÚCallableÚOptionalÚUnionN)ÚTensor)ÚLiteral)ÚMetric)Úrank_zero_warn)Údim_zero_cat)Ú_MATPLOTLIB_AVAILABLE)Ú_AX_TYPEÚ_PLOT_OUT_TYPE)ÚRunning)zSumMetric.plotzMeanMetric.plotzMaxMetric.plotzMinMetric.plotc                   óà   ‡ — e Zd ZU dZdZdZdZeed<   	 	 dde	e
ef   de	eef   de	ed	   ef   d
ededdfˆ fd„Z	 dde	eef   dee	eef      deeef   fd„Zde	eef   ddfd„Zdefd„Zˆ xZS )ÚBaseAggregatorae  Base class for aggregation metrics.

    Args:
        fn: string specifying the reduction function
        default_value: default tensor value to use for the metric state
        nan_strategy: options:
            - ``'error'``: if any `nan` values are encountered will give a RuntimeError
            - ``'warn'``: if any `nan` values are encountered will give a warning and continue
            - ``'ignore'``: all `nan` values are silently removed
            - ``'disable'``: disable all `nan` checks
            - a float: if a float is provided will impute any `nan` values with this value

        state_name: name of the metric state
        kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info.

    Raises:
        ValueError:
            If ``nan_strategy`` is not one of ``error``, ``warn``, ``ignore``, ``disable`` or a float

    NFÚfull_state_updateÚvalueÚfnÚdefault_valueÚnan_strategy©ÚerrorÚwarnÚignoreÚdisableÚ
state_nameÚkwargsÚreturnc                 ó¸   •— t        ‰| �  di |¤Ž d}||vr"t        |t        «      st	        d|› d|› d�«      ‚|| _        | j                  |||¬«       || _        y )Nr   z6Arg `nan_strategy` should either be a float or one of z	 but got ú.©ÚdefaultÚdist_reduce_fx© )ÚsuperÚ__init__Ú
isinstanceÚfloatÚ
ValueErrorr   Ú	add_stater   )Úselfr   r   r   r   r   Úallowed_nan_strategyÚ	__class__s          €úm/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchmetrics/aggregation.pyr'   zBaseAggregator.__init__:   sv   ø€ ô 	‰ÑÑ"˜6Ò"ØEÐØÐ3Ñ3¼JÀ|ÔUZÔ<[ÜØHÐI]ÐH^Ð^gÐhtÐguÐuvÐwóð ð )ˆÔØ�‰�z¨=ÈˆÔLØ$ˆ�ó    ÚxÚweightc                 ó  — t        |t        «      s,t        j                  || j                  | j
                  ¬«      }|�<t        |t        «      s,t        j                  || j                  | j
                  ¬«      }| j                  dk7  �r,t        j                  |«      }|�t        j                  |«      }n8t        j                  |«      j                  «       }t        j                  |«      }|j                  «       s|j                  «       r¼| j                  dk(  rt        d«      ‚| j                  dv r2| j                  dk(  rt        dt        «       |||z      }|||z      }nbt        | j                  t        «      st!        d| j                  › �«      ‚| j                  |||z  <   d	|||z  <   nt        j                  |«      }|j#                  | j                  «      |j#                  | j                  «      fS )
z3Convert input ``x`` to a tensor and check for Nans.©ÚdtypeÚdevicer   r   z"Encountered `nan` values in tensor)r   r   r   z4Encountered `nan` values in tensor. Will be removed.z+`nan_strategy` shall be float but you pass é   )r(   r   ÚtorchÚ	as_tensorr5   r6   r   ÚisnanÚ
zeros_likeÚboolÚ	ones_likeÚanyÚRuntimeErrorr   ÚUserWarningr)   r*   Úto)r,   r1   r2   ÚnansÚnans_weights        r/   Ú_cast_and_nan_check_inputz(BaseAggregator._cast_and_nan_check_inputM   s¢  € ô ˜!œVÔ$Ü—‘ ¨¯©¸D¿K¹KÔHˆAØÐ¤j°¼Ô&@Ü—_‘_ V°4·:±:ÀdÇkÁkÔRˆFà×Ñ 	Ó)Ü—;‘;˜q“>ˆDØÐ!Ü#Ÿk™k¨&Ó1‘ä#×.Ñ.¨tÓ4×9Ñ9Ó;�ÜŸ™¨Ó+�Ø�x‰xŒz˜[Ÿ_™_Ô.Ø×$Ñ$¨Ò/Ü&Ð'KÓLÐLØ×$Ñ$Ð(:Ñ:Ø×(Ñ(¨FÒ2Ü&Ð']Ô_jÔkØ˜D ;Ñ.Ð/Ñ0�AØ# d¨[Ñ&8Ð$9Ñ:‘Fä% d×&7Ñ&7¼Ô?Ü(Ð+VÐW[×WhÑWhÐViÐ)jÓkÐkØ,0×,=Ñ,=�A�d˜[Ñ(Ñ)Ø12�F˜4 +Ñ-Ò.ä—_‘_ QÓ'ˆFØ�t‰t�D—J‘JÓ §¡¨4¯:©:Ó!6Ð6Ð6r0   c                  ó   — y)zOverwrite in child class.Nr%   )r,   r   s     r/   ÚupdatezBaseAggregator.updaten   s   � r0   c                 ó.   — t        | | j                  «      S ©zCompute the aggregated value.)Úgetattrr   ©r,   s    r/   ÚcomputezBaseAggregator.computeq   s   € ä�t˜TŸ_™_Ó-Ð-r0   )r   r   ©N)Ú__name__Ú
__module__Ú__qualname__Ú__doc__Úis_differentiableÚhigher_is_betterr   r<   Ú__annotations__r   r   Ústrr   Úlistr	   r)   r   r'   r   ÚtuplerD   rF   rK   Ú__classcell__©r.   s   @r/   r   r       sù   ø… ñð* ÐØÐØ#Ð�tÓ#ð U\Ø!ñ%à�(˜C�-Ñ ð%ð ˜V T˜\Ñ*ð%ð ˜GÐ$HÑIÈ5ÐPÑQð	%ð
 ð%ð ð%ð 
õ%ð( QUñ7Ø�u˜f�}Ñ%ð7Ø/7¸¸eÀV¸mÑ8LÑ/Mð7à	ˆv�vˆ~Ñ	ó7ðB(˜E %¨ -Ñ0ð (°Tó (ð.˜÷ .r0   r   c                   ó¬   ‡ — e Zd ZU dZdZeed<   eed<   	 ddee	d   e
f   dedd	fˆ fd
„Zdee
ef   dd	fd„Z	 ddeeeee   f      dee   defd„Zˆ xZS )Ú	MaxMetrica}  Aggregate a stream of value into their maximum value.

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

    - ``value`` (:class:`~float` or :class:`~torch.Tensor`): a single float or an tensor of float values with
      arbitrary shape ``(...,)``.

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

    - ``agg`` (:class:`~torch.Tensor`): scalar float tensor with aggregated maximum value over all inputs received

    Args:
        nan_strategy: options:
            - ``'error'``: if any `nan` values are encountered will give a RuntimeError
            - ``'warn'``: if any `nan` values are encountered will give a warning and continue
            - ``'ignore'``: all `nan` values are silently removed
            - ``'disable'``: disable all `nan` checks
            - a float: if a float is provided will impute any `nan` values with this value

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

    Raises:
        ValueError:
            If ``nan_strategy`` is not one of ``error``, ``warn``, ``ignore``, ``disable`` or a float

    Example:
        >>> from torch import tensor
        >>> from torchmetrics.aggregation import MaxMetric
        >>> metric = MaxMetric()
        >>> metric.update(1)
        >>> metric.update(tensor([2, 3]))
        >>> metric.compute()
        tensor(3.)

    Tr   Ú	max_valuer   r   r   r   Nc                 ó�   •— t        ‰| �  dt        j                  t	        d«      t        j
                  «       ¬«       |fddi|¤Ž y )NÚmaxÚinf©r5   r   r[   ©r&   r'   r8   Útensorr)   Úget_default_dtype©r,   r   r   r.   s      €r/   r'   zMaxMetric.__init__ž   sJ   ø€ ô
 	‰ÑØÜ�\‰\œ% ›,¬e×.EÑ.EÓ.GÔHÐHØñ	
ð #ð		
ð
 ó	
r0   r   c                 ó¾   — | j                  |«      \  }}|j                  «       r9t        j                  | j                  t        j                  |«      «      | _        yy©z¬Update state with data.

        Args:
            value: Either a float or tensor containing data. Additional tensor
                dimensions will be flattened

        N)rD   Únumelr8   r]   r[   ©r,   r   Ú_s      r/   rF   zMaxMetric.update«   óE   € ð ×1Ñ1°%Ó8‰ˆˆqØ�;‰;Œ=Ü"ŸY™Y t§~¡~´u·y±yÀÓ7GÓHˆD�Nð 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

            >>> # Example plotting a single value
            >>> from torchmetrics.aggregation import MaxMetric
            >>> metric = MaxMetric()
            >>> metric.update([1, 2, 3])
            >>> fig_, ax_ = metric.plot()

        .. plot::
            :scale: 75

            >>> # Example plotting multiple values
            >>> from torchmetrics.aggregation import MaxMetric
            >>> metric = MaxMetric()
            >>> values = [ ]
            >>> for i in range(10):
            ...     values.append(metric(i))
            >>> fig_, ax_ = metric.plot(values)

        ©Ú_plot©r,   rj   rk   s      r/   ÚplotzMaxMetric.plot·   ó   € ðL �z‰z˜#˜rÓ"Ð"r0   ©r   ©NN©rM   rN   rO   rP   r   r<   rS   r   r   r	   r)   r   r'   rF   r   r   r   r   rp   rW   rX   s   @r/   rZ   rZ   v   ó±   ø… ñ"ðH #Ð�tÓ"ØÓð U[ñ
à˜GÐ$HÑIÈ5ÐPÑQð
ð ð
ð 
õ	
ð
I˜E %¨ -Ñ0ð 
I°Tó 
Ið _cñ&#Ø˜E &¨(°6Ñ*:Ð":Ñ;Ñ<ð&#ØIQÐRZÑI[ð&#à	÷&#r0   rZ   c                   ó¬   ‡ — e Zd ZU dZdZeed<   eed<   	 ddee	d   e
f   dedd	fˆ fd
„Zdee
ef   dd	fd„Z	 ddeeeee   f      dee   defd„Zˆ xZS )Ú	MinMetrica}  Aggregate a stream of value into their minimum value.

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

    - ``value`` (:class:`~float` or :class:`~torch.Tensor`): a single float or an tensor of float values with
      arbitrary shape ``(...,)``.

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

    - ``agg`` (:class:`~torch.Tensor`): scalar float tensor with aggregated minimum value over all inputs received

    Args:
        nan_strategy: options:
            - ``'error'``: if any `nan` values are encountered will give a RuntimeError
            - ``'warn'``: if any `nan` values are encountered will give a warning and continue
            - ``'ignore'``: all `nan` values are silently removed
            - ``'disable'``: disable all `nan` checks
            - a float: if a float is provided will impute any `nan` values with this value

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

    Raises:
        ValueError:
            If ``nan_strategy`` is not one of ``error``, ``warn``, ``ignore``, ``disable`` or a float

    Example:
        >>> from torch import tensor
        >>> from torchmetrics.aggregation import MinMetric
        >>> metric = MinMetric()
        >>> metric.update(1)
        >>> metric.update(tensor([2, 3]))
        >>> metric.compute()
        tensor(1.)

    Tr   Ú	min_valuer   r   r   r   Nc                 óŽ   •— t        ‰| �  dt        j                  t	        d«      t        j
                  «       ¬«      |fddi|¤Ž y )NÚminr^   r_   r   rx   r`   rc   s      €r/   r'   zMinMetric.__init__  sG   ø€ ô
 	‰ÑØÜ�L‰Lœ˜u›¬U×-DÑ-DÓ-FÔGØñ	
ð #ð		
ð
 ó	
r0   r   c                 ó¾   — | j                  |«      \  }}|j                  «       r9t        j                  | j                  t        j                  |«      «      | _        yyre   )rD   rf   r8   rz   rx   rg   s      r/   rF   zMinMetric.update  ri   r0   rj   rk   c                 ó&   — | 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
            >>> from torchmetrics.aggregation import MinMetric
            >>> metric = MinMetric()
            >>> metric.update([1, 2, 3])
            >>> fig_, ax_ = metric.plot()

        .. plot::
            :scale: 75

            >>> # Example plotting multiple values
            >>> from torchmetrics.aggregation import MinMetric
            >>> metric = MinMetric()
            >>> values = [ ]
            >>> for i in range(10):
            ...     values.append(metric(i))
            >>> fig_, ax_ = metric.plot(values)

        rm   ro   s      r/   rp   zMinMetric.plot!  rq   r0   rr   rs   rt   rX   s   @r/   rw   rw   à   ru   r0   rw   c                   óž   ‡ — e Zd ZU dZeed<   	 ddeed   ef   de	ddfˆ fd„Z
d	eeef   ddfd
„Z	 ddeeeee   f      dee   defd„Zˆ xZS )Ú	SumMetricai  Aggregate a stream of value into their sum.

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

    - ``value`` (:class:`~float` or :class:`~torch.Tensor`): a single float or an tensor of float values with
      arbitrary shape ``(...,)``.

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

    - ``agg`` (:class:`~torch.Tensor`): scalar float tensor with aggregated sum over all inputs received

    Args:
        nan_strategy: options:
            - ``'error'``: if any `nan` values are encountered will give a RuntimeError
            - ``'warn'``: if any `nan` values are encountered will give a warning and continue
            - ``'ignore'``: all `nan` values are silently removed
            - ``'disable'``: disable all `nan` checks
            - a float: if a float is provided will impute any `nan` values with this value

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

    Raises:
        ValueError:
            If ``nan_strategy`` is not one of ``error``, ``warn``, ``ignore``, ``disable`` or a float

    Example:
        >>> from torch import tensor
        >>> from torchmetrics.aggregation import SumMetric
        >>> metric = SumMetric()
        >>> metric.update(1)
        >>> metric.update(tensor([2, 3]))
        >>> metric.compute()
        tensor(6.)

    Ú	sum_valuer   r   r   r   Nc                 ó|   •— t        ‰| �  dt        j                  dt        j                  «       ¬«      |fddi|¤Ž y )NÚsumç        r_   r   r   )r&   r'   r8   ra   rb   rc   s      €r/   r'   zSumMetric.__init__q  sC   ø€ ô
 	‰ÑØÜ�L‰L˜¤E×$;Ñ$;Ó$=Ô>Øñ	
ð #ð		
ð
 ó	
r0   r   c                 ó”   — | j                  |«      \  }}|j                  «       r$| xj                  |j                  «       z  c_        yyre   )rD   rf   r   r�   rg   s      r/   rF   zSumMetric.update~  s:   € ð ×1Ñ1°%Ó8‰ˆˆqØ�;‰;Œ=Ø�NŠN˜eŸi™i›kÑ)ŽNð r0   rj   rk   c                 ó&   — | 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
            >>> from torchmetrics.aggregation import SumMetric
            >>> metric = SumMetric()
            >>> metric.update([1, 2, 3])
            >>> fig_, ax_ = metric.plot()

        .. plot::
            :scale: 75

            >>> # Example plotting multiple values
            >>> from torch import rand, randint
            >>> from torchmetrics.aggregation import SumMetric
            >>> metric = SumMetric()
            >>> values = [ ]
            >>> for i in range(10):
            ...     values.append(metric([i, i+1]))
            >>> fig_, ax_ = metric.plot(values)

        rm   ro   s      r/   rp   zSumMetric.plotŠ  s   € ðN �z‰z˜#˜rÓ"Ð"r0   rr   rs   )rM   rN   rO   rP   r   rS   r   r	   r)   r   r'   rF   r   r   r   r   rp   rW   rX   s   @r/   r~   r~   J  s£   ø… ñ"ðH Óð U[ñ
à˜GÐ$HÑIÈ5ÐPÑQð
ð ð
ð 
õ	
ð
*˜E %¨ -Ñ0ð 
*°Tó 
*ð _cñ'#Ø˜E &¨(°6Ñ*:Ð":Ñ;Ñ<ð'#ØIQÐRZÑI[ð'#à	÷'#r0   r~   c                   óv   ‡ — e Zd ZU dZeed<   	 ddeed   ef   de	ddfˆ fd„Z
deeef   ddfd	„Zdefd
„Zˆ xZS )Ú	CatMetricak  Concatenate a stream of values.

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

    - ``value`` (:class:`~float` or :class:`~torch.Tensor`): a single float or an tensor of float values with
      arbitrary shape ``(...,)``.

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

    - ``agg`` (:class:`~torch.Tensor`): scalar float tensor with concatenated values over all input received

    Args:
        nan_strategy: options:
            - ``'error'``: if any `nan` values are encountered will give a RuntimeError
            - ``'warn'``: if any `nan` values are encountered will give a warning and continue
            - ``'ignore'``: all `nan` values are silently removed
            - ``'disable'``: disable all `nan` checks
            - a float: if a float is provided will impute any `nan` values with this value

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

    Raises:
        ValueError:
            If ``nan_strategy`` is not one of ``error``, ``warn``, ``ignore``, ``disable`` or a float

    Example:
        >>> from torch import tensor
        >>> from torchmetrics.aggregation import CatMetric
        >>> metric = CatMetric()
        >>> metric.update(1)
        >>> metric.update(tensor([2, 3]))
        >>> metric.compute()
        tensor([1., 2., 3.])

    r   r   r   r   r   Nc                 ó*   •— t        ‰| �  dg |fi |¤Ž y )NÚcat)r&   r'   rc   s      €r/   r'   zCatMetric.__init__Û  s   ø€ ô
 	‰Ñ˜  LÑ;°FÓ;r0   c                 ó„   — | j                  |«      \  }}|j                  «       r| j                  j                  |«       yyre   )rD   rf   r   Úappendrg   s      r/   rF   zCatMetric.updateâ  s8   € ð ×1Ñ1°%Ó8‰ˆˆqØ�;‰;Œ=Ø�J‰J×Ñ˜eÕ$ð r0   c                 ó�   — t        | j                  t        «      r!| j                  rt        | j                  «      S | j                  S rH   )r(   r   rU   r   rJ   s    r/   rK   zCatMetric.computeî  s/   € ä�d—j‘j¤$Ô'¨D¯JªJÜ §
¡
Ó+Ð+Ø�z‰zÐr0   rr   )rM   rN   rO   rP   r   rS   r   r	   r)   r   r'   rF   rK   rW   rX   s   @r/   r†   r†   ´  sp   ø… ñ"ðH ƒMð U[ñ<à˜GÐ$HÑIÈ5ÐPÑQð<ð ð<ð 
õ	<ð
%˜E %¨ -Ñ0ð 
%°Tó 
%ð˜÷ r0   r†   c                   óÆ   ‡ — e Zd ZU dZeed<   eed<   	 ddeed   ef   de	ddfˆ fd	„Z
dd
eeef   deeedf   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 )Ú
MeanMetrica-  Aggregate a stream of value into their mean value.

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

    - ``value`` (:class:`~float` or :class:`~torch.Tensor`): a single float or an tensor of float values with
      arbitrary shape ``(...,)``.
    - ``weight`` (:class:`~float` or :class:`~torch.Tensor`): a single float or an tensor of float value with
      arbitrary shape ``(...,)``. Needs to be broadcastable with the shape of ``value`` tensor.

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

    - ``agg`` (:class:`~torch.Tensor`): scalar float tensor with aggregated (weighted) mean over all inputs received

    Args:
        nan_strategy: options:
            - ``'error'``: if any `nan` values are encountered will give a RuntimeError
            - ``'warn'``: if any `nan` values are encountered will give a warning and continue
            - ``'ignore'``: all `nan` values are silently removed
            - ``'disable'``: disable all `nan` checks
            - a float: if a float is provided will impute any `nan` values with this value

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

    Raises:
        ValueError:
            If ``nan_strategy`` is not one of ``error``, ``warn``, ``ignore``, ``disable`` or a float

    Example:
        >>> from torchmetrics.aggregation import MeanMetric
        >>> metric = MeanMetric()
        >>> metric.update(1)
        >>> metric.update(torch.tensor([2, 3]))
        >>> metric.compute()
        tensor(2.)

    Ú
mean_valuer2   r   r   r   r   Nc                 óò   •— t        ‰| �  dt        j                  dt        j                  «       ¬«      |fddi|¤Ž | j                  dt        j                  dt        j                  «       ¬«      d¬«       y )Nr�   r‚   r_   r   rŽ   r2   r"   )r&   r'   r8   ra   rb   r+   rc   s      €r/   r'   zMeanMetric.__init__  sl   ø€ ô
 	‰ÑØÜ�L‰L˜¤E×$;Ñ$;Ó$=Ô>Øñ	
ð $ð		
ð
 ò	
ð 	�‰�x¬¯©°cÄ×AXÑAXÓAZÔ)[ÐlqˆÕrr0   r   c                 óH  — t        |t        «      s,t        j                  || j                  | j
                  ¬«      }|€t        j                  |«      }n<t        |t        «      s,t        j                  || j                  | j
                  ¬«      }t        j                  ||j                  «      }| j                  ||«      \  }}|j                  «       dk(  ry| xj                  ||z  j                  «       z  c_        | xj                  |j                  «       z  c_        y)aº  Update state with data.

        Args:
            value: Either a float or tensor containing data. Additional tensor
                dimensions will be flattened
            weight: Either a float or tensor containing weights for calculating
                the average. Shape of weight should be able to broadcast with
                the shape of `value`. Default to None corresponding to simple
                harmonic average.

        r4   Nr   )r(   r   r8   r9   r5   r6   r=   Úbroadcast_toÚshaperD   rf   rŽ   r�   r2   )r,   r   r2   s      r/   rF   zMeanMetric.update,  sË   € ô ˜%¤Ô(Ü—O‘O E°·±ÀDÇKÁKÔPˆEØˆ>Ü—_‘_ UÓ+‰FÜ˜F¤FÔ+Ü—_‘_ V°4·:±:ÀdÇkÁkÔRˆFÜ×#Ñ# F¨E¯K©KÓ8ˆØ×6Ñ6°u¸fÓE‰ˆˆvà�;‰;‹=˜AÒØØ�Š˜E F™N×/Ñ/Ó1Ñ1�Ø�Š�v—z‘z“|Ñ#Žr0   c                 ó4   — | j                   | j                  z  S rH   )rŽ   r2   rJ   s    r/   rK   zMeanMetric.computeG  s   € à�‰ §¡Ñ,Ð,r0   rj   rk   c                 ó&   — | 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
            >>> from torchmetrics.aggregation import MeanMetric
            >>> metric = MeanMetric()
            >>> metric.update([1, 2, 3])
            >>> fig_, ax_ = metric.plot()

        .. plot::
            :scale: 75

            >>> # Example plotting multiple values
            >>> from torchmetrics.aggregation import MeanMetric
            >>> metric = MeanMetric()
            >>> values = [ ]
            >>> for i in range(10):
            ...     values.append(metric([i, i+1]))
            >>> fig_, ax_ = metric.plot(values)

        rm   ro   s      r/   rp   zMeanMetric.plotK  rq   r0   rr   rL   rs   )rM   rN   rO   rP   r   rS   r   r	   r)   r   r'   rF   rK   r   r   r   r   rp   rW   rX   s   @r/   r�   r�   õ  sÌ   ø… ñ#ðJ ÓØƒNð U[ñsà˜GÐ$HÑIÈ5ÐPÑQðsð ðsð 
õ	sñ$˜E %¨ -Ñ0ð $¸%ÀÀvÈtÐ@SÑ:Tð $Ð`dó $ð6-˜ó -ð
 _cñ&#Ø˜E &¨(°6Ñ*:Ð":Ñ;Ñ<ð&#ØIQÐRZÑI[ð&#à	÷&#r0   r�   c            	       óJ   ‡ — e Zd ZdZ	 	 d	dedeed   ef   deddfˆ fd„Z	ˆ xZ
S )
ÚRunningMeanai	  Aggregate a stream of value into their mean over a running window.

    Using this metric compared to `MeanMetric` allows for calculating metrics over a running window of values, instead
    of the whole history of values. This is beneficial when you want to get a better estimate of the metric during
    training and don't want to wait for the whole training to finish to get epoch level estimates.

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

    - ``value`` (:class:`~float` or :class:`~torch.Tensor`): a single float or an tensor of float values with
      arbitrary shape ``(...,)``.

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

    - ``agg`` (:class:`~torch.Tensor`): scalar float tensor with aggregated sum over all inputs received

    Args:
        nan_strategy: options:
            - ``'error'``: if any `nan` values are encountered will give a RuntimeError
            - ``'warn'``: if any `nan` values are encountered will give a warning and continue
            - ``'ignore'``: all `nan` values are silently removed
            - ``'disable'``: disable all `nan` checks
            - a float: if a float is provided will impute any `nan` values with this value

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

    Raises:
        ValueError:
            If ``nan_strategy`` is not one of ``error``, ``warn``, ``ignore``, ``disable`` or a float

    Example:
        >>> from torch import tensor
        >>> from torchmetrics.aggregation import RunningMean
        >>> metric = RunningMean(window=3)
        >>> for i in range(6):
        ...     current_val = metric(tensor([i]))
        ...     running_val = metric.compute()
        ...     total_val = tensor(sum(list(range(i+1)))) / (i+1)  # total mean over all samples
        ...     print(f"{current_val=}, {running_val=}, {total_val=}")
        current_val=tensor(0.), running_val=tensor(0.), total_val=tensor(0.)
        current_val=tensor(1.), running_val=tensor(0.5000), total_val=tensor(0.5000)
        current_val=tensor(2.), running_val=tensor(1.), total_val=tensor(1.)
        current_val=tensor(3.), running_val=tensor(2.), total_val=tensor(1.5000)
        current_val=tensor(4.), running_val=tensor(3.), total_val=tensor(2.)
        current_val=tensor(5.), running_val=tensor(4.), total_val=tensor(2.5000)

    Úwindowr   r   r   r   Nc                 ó>   •— t         ‰| �  t        dd|i|¤Ž|¬«       y ©Nr   )Úbase_metricr—   r%   )r&   r'   r�   ©r,   r—   r   r   r.   s       €r/   r'   zRunningMean.__init__¤  s%   ø€ ô 	‰Ñ¤ZÑ%T¸\Ð%TÈVÑ%TÐ]cÐÕdr0   ©é   r   ©rM   rN   rO   rP   Úintr   r	   r)   r   r'   rW   rX   s   @r/   r–   r–   t  sW   ø„ ñ-ðb ØTZñeàðeð ˜GÐ$HÑIÈ5ÐPÑQðeð ð	eð
 
÷eñ er0   r–   c            	       óJ   ‡ — e Zd ZdZ	 	 d	dedeed   ef   deddfˆ fd„Z	ˆ xZ
S )
Ú
RunningSumay	  Aggregate a stream of value into their sum over a running window.

    Using this metric compared to `SumMetric` allows for calculating metrics over a running window of values, instead
    of the whole history of values. This is beneficial when you want to get a better estimate of the metric during
    training and don't want to wait for the whole training to finish to get epoch level estimates.

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

    - ``value`` (:class:`~float` or :class:`~torch.Tensor`): a single float or an tensor of float values with
      arbitrary shape ``(...,)``.

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

    - ``agg`` (:class:`~torch.Tensor`): scalar float tensor with aggregated sum over all inputs received

    Args:
        window: The size of the running window.
        nan_strategy: options:
            - ``'error'``: if any `nan` values are encountered will give a RuntimeError
            - ``'warn'``: if any `nan` values are encountered will give a warning and continue
            - ``'ignore'``: all `nan` values are silently removed
            - ``'disable'``: disable all `nan` checks
            - a float: if a float is provided will impute any `nan` values with this value

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

    Raises:
        ValueError:
            If ``nan_strategy`` is not one of ``error``, ``warn``, ``ignore``, ``disable`` or a float

    Example:
        >>> from torch import tensor
        >>> from torchmetrics.aggregation import RunningSum
        >>> metric = RunningSum(window=3)
        >>> for i in range(6):
        ...     current_val = metric(tensor([i]))
        ...     running_val = metric.compute()
        ...     total_val = tensor(sum(list(range(i+1))))  # total sum over all samples
        ...     print(f"{current_val=}, {running_val=}, {total_val=}")
        current_val=tensor(0.), running_val=tensor(0.), total_val=tensor(0)
        current_val=tensor(1.), running_val=tensor(1.), total_val=tensor(1)
        current_val=tensor(2.), running_val=tensor(3.), total_val=tensor(3)
        current_val=tensor(3.), running_val=tensor(6.), total_val=tensor(6)
        current_val=tensor(4.), running_val=tensor(9.), total_val=tensor(10)
        current_val=tensor(5.), running_val=tensor(12.), total_val=tensor(15)

    r—   r   r   r   r   Nc                 ó>   •— t         ‰| �  t        dd|i|¤Ž|¬«       y r™   )r&   r'   r~   r›   s       €r/   r'   zRunningSum.__init__Þ  s%   ø€ ô 	‰Ñ¤YÑ%S¸LÐ%SÈFÑ%SÐ\bÐÕcr0   rœ   rž   rX   s   @r/   r¡   r¡   ­  sW   ø„ ñ.ðd ØTZñdàðdð ˜GÐ$HÑIÈ5ÐPÑQðdð ð	dð
 
÷dñ dr0   r¡   )!Úcollections.abcr   Útypingr   r   r   r   r8   r   Útyping_extensionsr	   Útorchmetrics.metricr
   Útorchmetrics.utilitiesr   Útorchmetrics.utilities.datar   Útorchmetrics.utilities.importsr   Útorchmetrics.utilities.plotr   r   Útorchmetrics.wrappers.runningr   Ú__doctest_skip__r   rZ   rw   r~   r†   r�   r–   r¡   r%   r0   r/   ú<module>r­      s©   ðõ %ß 1Ó 1ã Ý Ý %å &Ý 1Ý 4Ý @ß @Ý 1áÚ`ÐôS.�Vô S.ôlg#�ô g#ôTg#�ô g#ôTg#�ô g#ôT>�ô >ôB|#�ô |#ô~6e�'ô 6eôr7d�õ 7dr0   