Ë
    þÍ:js#  ã                   óà   — d dl mZ 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mZ d dlmZ esdgZ	 ddedede	j                  fd„Z G d„ de«      Zy)é    )ÚSequence)Údeepcopy)ÚAnyÚOptionalÚUnionÚcastN)Úapply_to_collection)ÚTensor)Ú
ModuleList)ÚMetric)Ú_MATPLOTLIB_AVAILABLE)Ú_AX_TYPEÚ_PLOT_OUT_TYPE)ÚWrapperMetriczBootStrapper.plotÚsizeÚsampling_strategyÚreturnc                 óL  — |dk(  ret         j                  j                  d«      }|j                  | f«      }t        j                  | «      j                  |j                  «       d¬«      S |dk(  r+t        j                  t        j                  | «      | d¬«      S t        d«      ‚)	züResample a tensor along its first dimension with replacement.

    Args:
        size: number of samples
        sampling_strategy: the strategy to use for sampling, either ``'poisson'`` or ``'multinomial'``

    Returns:
        resampled tensor

    Úpoissoné   r   ©ÚdimÚmultinomialT)Únum_samplesÚreplacementzUnknown sampling strategy)
ÚtorchÚdistributionsÚPoissonÚsampleÚarangeÚrepeat_interleaveÚlongr   ÚonesÚ
ValueError)r   r   ÚpÚns       úx/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchmetrics/wrappers/bootstrapping.pyÚ_bootstrap_samplerr(       sŠ   € ð ˜IÒ%Ü×Ñ×'Ñ'¨Ó*ˆØ�H‰H�d�WÓˆÜ�|‰|˜DÓ!×3Ñ3°A·F±F³HÀ!Ð3ÓDÐDØ˜MÒ)Ü× Ñ ¤§¡¨DÓ!1¸tÐQUÔVÐVÜ
Ð0Ó
1Ð1ó    c                   óþ   ‡ — e Zd ZU dZdZee   ed<   	 	 	 	 	 	 ddede	deded	ee
eef      d
e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	 ddee
eee   f      dee   defd„Zˆ xZS )ÚBootStrapperaÑ  Using `Turn a Metric into a Bootstrapped`_.

    That can automate the process of getting confidence intervals for metric values. This wrapper
    class basically keeps multiple copies of the same base metric in memory and whenever ``update`` or
    ``forward`` is called, all input tensors are resampled (with replacement) along the first dimension.

    Args:
        base_metric: base metric class to wrap
        num_bootstraps: number of copies to make of the base metric for bootstrapping
        mean: if ``True`` return the mean of the bootstraps
        std: if ``True`` return the standard deviation of the bootstraps
        quantile: if given, returns the quantile of the bootstraps. Can only be used with pytorch version 1.6 or higher
        raw: if ``True``, return all bootstrapped values
        sampling_strategy:
            Determines how to produce bootstrapped samplings. Either ``'poisson'`` or ``multinomial``.
            If ``'possion'`` is chosen, the number of times each sample will be included in the bootstrap
            will be given by :math:`n\sim Poisson(\lambda=1)`, which approximates the true bootstrap distribution
            when the number of samples is large. If ``'multinomial'`` is chosen, we will apply true bootstrapping
            at the batch level to approximate bootstrapping over the hole dataset.
        kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info.

    Example::
        >>> from pprint import pprint
        >>> from torch import randint
        >>> from torchmetrics.wrappers import BootStrapper
        >>> from torchmetrics.classification import MulticlassAccuracy
        >>> base_metric = MulticlassAccuracy(num_classes=5, average='micro')
        >>> bootstrap = BootStrapper(base_metric, num_bootstraps=20)
        >>> bootstrap.update(randint(5, (20,)), randint(5, (20,)))
        >>> output = bootstrap.compute()
        >>> pprint(output)
        {'mean': tensor(0.2089), 'std': tensor(0.0772)}

    TÚfull_state_updateNÚbase_metricÚnum_bootstrapsÚmeanÚstdÚquantileÚrawr   Úkwargsr   c                 óL  •— t        ‰| �  di |¤Ž t        |t        «      st	        d|› �«      ‚t        t        |«      D �	cg c]  }	t        |«      ‘Œ c}	«      | _        || _	        || _
        || _        || _        || _        d}
||
vrt	        d|
› d|› �«      ‚|| _        y c c}	w )NzKExpected base metric to be an instance of torchmetrics.Metric but received )r   r   z5Expected argument ``sampling_strategy`` to be one of z but received © )ÚsuperÚ__init__Ú
isinstancer   r$   r   Úranger   Úmetricsr.   r/   r0   r1   r2   r   )Úselfr-   r.   r/   r0   r1   r2   r   r3   Ú_Úallowed_samplingÚ	__class__s              €r'   r7   zBootStrapper.__init__]   sÂ   ø€ ô 	‰ÑÑ"˜6Ò"Ü˜+¤vÔ.ÜØ]Ð^iÐ]jÐkóð ô "Ä%ÈÓBWÖ"X¸Q¤8¨KÕ#8Ò"XÓYˆŒØ,ˆÔàˆŒ	ØˆŒØ ˆŒØˆŒà5ÐØÐ$4Ñ4ÜØGÐHXÐGYØ Ð!2Ð 3ð5óð ð "3ˆÕùò #Ys   ÁB!Úargsc                 óÊ  — t        |t        j                  t        «      }t        |t        j                  t        «      }t        |«      dkD  r|d   }n<t        |«      dkD  r#t	        t        |j                  «       «      «      }nt        d«      ‚t        | j                  «      D ]½  }t        || j                  ¬«      j                  | j                  «      }|j                  «       dk(  rŒGt        |t        j                  t        j                  d|¬«      }t        |t        j                  t        j                  d|¬«      }	 | j                   |   j"                  |i |	¤Ž Œ¿ y)ztUpdate the state of the base metric.

        Any tensor passed in will be bootstrapped along dimension 0.

        r   zMNone of the input contained tensors, so could not determine the sampling size)r   )r   ÚindexN)r	   r   r
   ÚlenÚnextÚiterÚvaluesr$   r9   r.   r(   r   ÚtoÚdeviceÚnumelÚindex_selectr:   Úupdate)
r;   r?   r3   Ú
args_sizesÚkwargs_sizesr   ÚidxÚ
sample_idxÚnew_argsÚ
new_kwargss
             r'   rJ   zBootStrapper.update~   s  € ô )¨¬u¯|©|¼SÓAˆ
Ü*¨6´5·<±<ÄÓEˆÜˆz‹?˜QÒØ˜a‘=‰DÜ�Ó Ò"Üœ˜\×0Ñ0Ó2Ó3Ó4‰DäÐlÓmÐmä˜×,Ñ,Ó-ò 	>ˆCÜ+¨DÀD×DZÑDZÔ[×^Ñ^Ð_c×_jÑ_jÓkˆJØ×ÑÓ! QÒ&ØÜ*¨4´·±¼u×?QÑ?QÐWXÐ`jÔkˆHÜ,¨V´U·\±\Ä5×CUÑCUÐ[\ÐdnÔoˆJØ$ˆD�L‰L˜Ñ×$Ñ$ hÐ=°*Ó=ñ	>r)   c           	      ó®  — t        j                  | j                  D �cg c]   }t        t        |«      j                  «       ‘Œ" c}d¬«      }i }| j                  r|j                  d¬«      |d<   | j                  r|j                  d¬«      |d<   | j                  �#t        j                  || j                  «      |d<   | j                  r||d<   |S c c}w )zâCompute the bootstrapped metric values.

        Always returns a dict of tensors, which can contain the following keys: ``mean``, ``std``, ``quantile`` and
        ``raw`` depending on how the class was initialized.

        r   r   r/   r0   r1   r2   )
r   Ústackr:   r   r   Úcomputer/   r0   r1   r2   )r;   ÚmÚcomputed_valsÚoutput_dicts       r'   rS   zBootStrapper.compute•   s¶   € ô Ÿ™ÈÏÉÖ$UÀ1¤T¬&°!£_×%<Ñ%<Õ%>Ò$UÐ[\Ô]ˆØˆØ�9Š9Ø"/×"4Ñ"4¸Ð"4Ó";ˆK˜ÑØ�8Š8Ø!.×!2Ñ!2°qÐ!2Ó!9ˆK˜ÑØ�=‰=Ð$Ü&+§n¡n°]ÀDÇMÁMÓ&RˆK˜
Ñ#Ø�8Š8Ø!.ˆK˜ÑØÐùò %Vs   ž%Cc                 ó*   •— t        t        | �
  |i |¤ŽS )z9Use the original forward method of the base metric class.)r6   r   Úforward)r;   r?   r3   r>   s      €r'   rX   zBootStrapper.forward¨   s   ø€ ä”] DÑ1°4ÐB¸6ÑBÐBr)   c                 ó„   •— | j                   D ]"  }t        t        |«      }|j                  «        Œ$ t        ‰| �  «        y)z#Reset the state of the base metric.N)r:   r   r   Úresetr6   )r;   rT   r>   s     €r'   rZ   zBootStrapper.reset¬   s5   ø€ à—‘ò 	ˆAÜ”V˜Q“ˆAØ�G‰G�Ið	ô 	‰‰�r)   Ú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
            >>> import torch
            >>> from torchmetrics.wrappers import BootStrapper
            >>> from torchmetrics.regression import MeanSquaredError
            >>> metric = BootStrapper(MeanSquaredError(), num_bootstraps=20)
            >>> metric.update(torch.randn(100,), torch.randn(100,))
            >>> fig_, ax_ = metric.plot()

        .. plot::
            :scale: 75

            >>> # Example plotting multiple values
            >>> import torch
            >>> from torchmetrics.wrappers import BootStrapper
            >>> from torchmetrics.regression import MeanSquaredError
            >>> metric = BootStrapper(MeanSquaredError(), num_bootstraps=20)
            >>> values = [ ]
            >>> for _ in range(3):
            ...     values.append(metric(torch.randn(100,), torch.randn(100,)))
            >>> fig_, ax_ = metric.plot(values)

        )Ú_plot)r;   r[   r\   s      r'   ÚplotzBootStrapper.plot³   s   € ðT �z‰z˜#˜rÓ"Ð"r)   )é
   TTNFr   )r   N)NN)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r,   r   ÚboolÚ__annotations__r   Úintr   Úfloatr
   Ústrr   r7   rJ   ÚdictrS   rX   rZ   r   r   r   r_   Ú__classcell__)r>   s   @r'   r+   r+   7   s*  ø… ñ!ðF )-Ð�x ‘~Ó,ð
 !ØØØ37ØØ!*ñ3àð3ð ð3ð ð	3ð
 ð3ð ˜5 ¨ Ñ/Ñ0ð3ð ð3ð ð3ð ð3ð 
õ3ðB>˜Cð >¨3ð >°4ó >ð.˜˜c 6˜kÑ*ó ð&C˜Sð C¨Cð C°Cõ Cõð _cñ*#Ø˜E &¨(°6Ñ*:Ð":Ñ;Ñ<ð*#ØIQÐRZÑI[ð*#à	÷*#r)   r+   )r   )Úcollections.abcr   Úcopyr   Útypingr   r   r   r   r   Úlightning_utilitiesr	   r
   Útorch.nnr   Útorchmetrics.metricr   Útorchmetrics.utilities.importsr   Útorchmetrics.utilities.plotr   r   Útorchmetrics.wrappers.abstractr   Ú__doctest_skip__rg   ri   r(   r+   r5   r)   r'   ú<module>rv      sm   ðõ %Ý ß -Ó -ã Ý 3Ý Ý å &Ý @ß @Ý 8áØ+Ð,Ðð
 'ñ2Ø
ð2àð2ð ‡\�\ó2ô.f#�=õ f#r)   