Ë
    þÍ:jv  ã                   óŒ   — 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mZ d dlmZ esd	gZ G d
„ de«      Zy)é    )ÚSequence)ÚAnyÚOptionalÚUnionN)ÚTensor)ÚMetric)Ú_MATPLOTLIB_AVAILABLE)Ú_AX_TYPEÚ_PLOT_OUT_TYPE)ÚWrapperMetriczMinMaxMetric.plotc                   ó   ‡ — e Zd ZU dZdZee   ed<   eed<   eed<   de	de
dd	fˆ fd
„Zde
de
dd	fd„Zdeeef   fd„Zde
de
de
fˆ fd„Zdˆ fd„Zedeeef   defd„«       Z	 ddeeeee   f      dee   defd„Zˆ xZS )ÚMinMaxMetricak  Wrapper metric that tracks both the minimum and maximum of a scalar/tensor across an experiment.

    The min/max value will be updated each time ``.compute`` is called.

    Args:
        base_metric:
            The metric of which you want to keep track of its maximum and minimum values.
        kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info.

    Raises:
        ValueError
            If ``base_metric` argument is not a subclasses instance of ``torchmetrics.Metric``

    Example::
        >>> import torch
        >>> from torchmetrics.wrappers import MinMaxMetric
        >>> from torchmetrics.classification import BinaryAccuracy
        >>> from pprint import pprint
        >>> base_metric = BinaryAccuracy()
        >>> minmax_metric = MinMaxMetric(base_metric)
        >>> preds_1 = torch.Tensor([[0.1, 0.9], [0.2, 0.8]])
        >>> preds_2 = torch.Tensor([[0.9, 0.1], [0.2, 0.8]])
        >>> labels = torch.Tensor([[0, 1], [0, 1]]).long()
        >>> pprint(minmax_metric(preds_1, labels))
        {'max': tensor(1.), 'min': tensor(1.), 'raw': tensor(1.)}
        >>> pprint(minmax_metric.compute())
        {'max': tensor(1.), 'min': tensor(1.), 'raw': tensor(1.)}
        >>> minmax_metric.update(preds_2, labels)
        >>> pprint(minmax_metric.compute())
        {'max': tensor(1.), 'min': tensor(0.7500), 'raw': tensor(0.7500)}

    TÚfull_state_updateÚmin_valÚmax_valÚbase_metricÚkwargsÚreturnNc                 óú   •— t        ‰| �  di |¤Ž t        |t        «      st	        d|› �«      ‚|| _        t        j                  t        d«      «      | _	        t        j                  t        d«      «      | _
        y )NzMExpected base metric to be an instance of `torchmetrics.Metric` but received Úinfz-inf© )ÚsuperÚ__init__Ú
isinstancer   Ú
ValueErrorÚ_base_metricÚtorchÚtensorÚfloatr   r   )Úselfr   r   Ú	__class__s      €úq/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchmetrics/wrappers/minmax.pyr   zMinMaxMetric.__init__D   sj   ø€ ô
 	‰ÑÑ"˜6Ò"Ü˜+¤vÔ.ÜØ_Ð`kÐ_lÐmóð ð (ˆÔÜ—|‘|¤E¨%£LÓ1ˆŒÜ—|‘|¤E¨&£MÓ2ˆ�ó    Úargsc                 ó<   —  | j                   j                  |i |¤Ž y)zUpdate the underlying metric.N)r   Úupdate)r    r$   r   s      r"   r&   zMinMaxMetric.updateR   s   € à ˆ×Ñ× Ñ  $Ð1¨&Ó1r#   c                 óú  — | j                   j                  «       }| j                  |«      st        d|› d�«      ‚| j                  j                  |j                  «      |k  r|n$| j                  j                  |j                  «      | _        | j                  j                  |j                  «      |kD  r|n$| j                  j                  |j                  «      | _        || j                  | j                  dœS )zêCompute the underlying metric as well as max and min values for this metric.

        Returns a dictionary that consists of the computed value (``raw``), as well as the minimum (``min``) and maximum
        (``max``) values.

        zLReturned value from base metric should be a float or scalar tensor, but got ú.)ÚrawÚmaxÚmin)r   ÚcomputeÚ_is_suitable_valÚRuntimeErrorr   ÚtoÚdevicer   )r    Úvals     r"   r,   zMinMaxMetric.computeV   s½   € ð ×Ñ×'Ñ'Ó)ˆØ×$Ñ$ SÔ)ÜÐ!mÐnqÐmrÐrsÐtÓuÐuØ"Ÿl™lŸo™o¨c¯j©jÓ9¸CÒ?‘sÀTÇ\Á\Ç_Á_ÐUX×U_ÑU_ÓE`ˆŒØ"Ÿl™lŸo™o¨c¯j©jÓ9¸CÒ?‘sÀTÇ\Á\Ç_Á_ÐUX×U_ÑU_ÓE`ˆŒØ 4§<¡<¸¿¹ÑEÐEr#   c                 ó*   •— t        t        | �
  |i |¤ŽS )z9Use the original forward method of the base metric class.)r   r   Úforward)r    r$   r   r!   s      €r"   r3   zMinMaxMetric.forwardd   s   ø€ ä”] DÑ1°4ÐB¸6ÑBÐBr#   c                 óV   •— t         ‰| �  «        | j                  j                  «        y)zXSet ``max_val`` and ``min_val`` to the initialization bounds and resets the base metric.N)r   Úresetr   )r    r!   s    €r"   r5   zMinMaxMetric.reseth   s   ø€ ä‰‰ŒØ×Ñ×ÑÕ!r#   r1   c                 óx   — t        | t        t        f«      ryt        | t        «      r| j	                  «       dk(  S y)z(Check whether min/max is a scalar value.Té   F)r   Úintr   r   Únumel)r1   s    r"   r-   zMinMaxMetric._is_suitable_valm   s3   € ô �cœC¤˜<Ô(ØÜ�cœ6Ô"Ø—9‘9“; !Ñ#Ð#Ør#   Ú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
            >>> import torch
            >>> from torchmetrics.wrappers import MinMaxMetric
            >>> from torchmetrics.classification import BinaryAccuracy
            >>> metric = MinMaxMetric(BinaryAccuracy())
            >>> metric.update(torch.randint(2, (20,)), torch.randint(2, (20,)))
            >>> fig_, ax_ = metric.plot()

        .. plot::
            :scale: 75

            >>> # Example plotting multiple values
            >>> import torch
            >>> from torchmetrics.wrappers import MinMaxMetric
            >>> from torchmetrics.classification import BinaryAccuracy
            >>> metric = MinMaxMetric(BinaryAccuracy())
            >>> values = [ ]
            >>> for _ in range(3):
            ...     values.append(metric(torch.randint(2, (20,)), torch.randint(2, (20,))))
            >>> fig_, ax_ = metric.plot(values)

        )Ú_plot)r    r1   r:   s      r"   ÚplotzMinMaxMetric.plotv   s   € ðT �z‰z˜#˜rÓ"Ð"r#   )r   N)NN)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   r   ÚboolÚ__annotations__r   r   r   r   r&   ÚdictÚstrr,   r3   r5   Ústaticmethodr   r   r-   r   r
   r   r=   Ú__classcell__)r!   s   @r"   r   r      s  ø… ñðB )-Ð�x ‘~Ó,ØƒOØƒOð3àð3ð ð3ð 
õ	3ð2˜Cð 2¨3ð 2°4ó 2ðF˜˜c 6˜kÑ*ó FðC˜Sð C¨Cð C°Cõ Cõ"ð
 ð˜e E¨6 MÑ2ð °tò ó ðð _cñ*#Ø˜E &¨(°6Ñ*:Ð":Ñ;Ñ<ð*#ØIQÐRZÑI[ð*#à	÷*#r#   r   )Úcollections.abcr   Útypingr   r   r   r   r   Útorchmetrics.metricr   Útorchmetrics.utilities.importsr	   Útorchmetrics.utilities.plotr
   r   Útorchmetrics.wrappers.abstractr   Ú__doctest_skip__r   r   r#   r"   ú<module>rO      s<   ðõ %ß 'Ñ 'ã Ý å &Ý @ß @Ý 8áØ+Ð,ÐôB#�=õ B#r#   