Ë
    þÍ:j�1  ã                   ó¾   — d dl mZ d dlmZmZmZ d dlmZmZ d dl	m
Z
mZmZ d dlmZ d dlmZ d dlmZmZ esg d¢Z G d	„ d
e«      Z G d„ de«      Z G d„ de«      Zy)é    )ÚSequence)ÚAnyÚOptionalÚUnion)ÚTensorÚtensor)Ú*complex_scale_invariant_signal_noise_ratioÚ"scale_invariant_signal_noise_ratioÚsignal_noise_ratio)ÚMetric)Ú_MATPLOTLIB_AVAILABLE)Ú_AX_TYPEÚ_PLOT_OUT_TYPE)zSignalNoiseRatio.plotz#ScaleInvariantSignalNoiseRatio.plotz*ComplexScaleInvariantSignalNoiseRatio.plotc                   óð   ‡ — e Zd ZU dZdZeed<   dZeed<   dZeed<   e	ed<   e	ed<   d	Z
ee   ed
<   d	Ze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 )ÚSignalNoiseRatioa3  Calculate `Signal-to-noise ratio`_ (SNR_) meric for evaluating quality of audio.

    .. math::
        \text{SNR} = \frac{P_{signal}}{P_{noise}}

    where  :math:`P` denotes the power of each signal. The SNR metric compares the level of the desired signal to
    the level of background noise. Therefore, a high value of SNR means that the audio is clear.

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

    - ``preds`` (:class:`~torch.Tensor`): float tensor with shape ``(...,time)``
    - ``target`` (:class:`~torch.Tensor`): float tensor with shape ``(...,time)``

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

    - ``snr`` (:class:`~torch.Tensor`): float scalar tensor with average SNR value over samples

    Args:
        zero_mean: if to zero mean target and preds or not
        kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info.

    Raises:
        TypeError:
            if target and preds have a different shape

    Example:
        >>> from torch import tensor
        >>> from torchmetrics.audio import SignalNoiseRatio
        >>> target = tensor([3.0, -0.5, 2.0, 7.0])
        >>> preds = tensor([2.5, 0.0, 2.0, 8.0])
        >>> snr = SignalNoiseRatio()
        >>> snr(preds, target)
        tensor(16.1805)

    FÚfull_state_updateTÚis_differentiableÚhigher_is_betterÚsum_snrÚtotalNÚplot_lower_boundÚplot_upper_boundÚ	zero_meanÚkwargsÚreturnc                 ó¦   •— t        ‰| �  di |¤Ž || _        | j                  dt	        d«      d¬«       | j                  dt	        d«      d¬«       y )Nr   ç        Úsum©ÚdefaultÚdist_reduce_fxr   r   © )ÚsuperÚ__init__r   Ú	add_stater   ©Úselfr   r   Ú	__class__s      €úk/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchmetrics/audio/snr.pyr$   zSignalNoiseRatio.__init__Q   sH   ø€ ô
 	‰ÑÑ"˜6Ò"Ø"ˆŒà�‰�y¬&°«+ÀeˆÔLØ�‰�w¬¨q«	À%ˆÕHó    ÚpredsÚtargetc                 óÀ   — t        ||| j                  ¬«      }| xj                  |j                  «       z  c_        | xj                  |j                  «       z  c_        y©ú*Update state with predictions and targets.)r+   r,   r   N)r   r   r   r   r   Únumel)r'   r+   r,   Ú	snr_batchs       r)   ÚupdatezSignalNoiseRatio.update\   s=   € ä&¨U¸6ÈTÏ^É^Ô\ˆ	à�Š˜	Ÿ™›Ñ'�Ø�
Š
�i—o‘oÓ'Ñ'Ž
r*   c                 ó4   — | j                   | j                  z  S ©zCompute metric.)r   r   ©r'   s    r)   ÚcomputezSignalNoiseRatio.computec   s   € à�|‰|˜dŸj™jÑ(Ð(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.audio import SignalNoiseRatio
            >>> metric = SignalNoiseRatio()
            >>> metric.update(torch.rand(4), torch.rand(4))
            >>> fig_, ax_ = metric.plot()

        .. plot::
            :scale: 75

            >>> # Example plotting multiple values
            >>> import torch
            >>> from torchmetrics.audio import SignalNoiseRatio
            >>> metric = SignalNoiseRatio()
            >>> values = [ ]
            >>> for _ in range(10):
            ...     values.append(metric(torch.rand(4), torch.rand(4)))
            >>> fig_, ax_ = metric.plot(values)

        ©Ú_plot©r'   r7   r8   s      r)   ÚplotzSignalNoiseRatio.plotg   s   € ðP �z‰z˜#˜rÓ"Ð"r*   ©F©NN)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   ÚboolÚ__annotations__r   r   r   r   r   Úfloatr   r   r$   r2   r6   r   r   r   r   r=   Ú__classcell__©r(   s   @r)   r   r   $   sâ   ø… ñ"ðH $Ð�tÓ#Ø"Ð�tÓ"Ø!Ð�dÓ!ØƒOØƒMØ(,Ð�h˜u‘oÓ,Ø(,Ð�h˜u‘oÓ,ð  ñ	Iàð	Ið ð	Ið 
õ		Ið(˜Fð (¨Fð (°tó (ð)˜ó )ð
 _cñ(#Ø˜E &¨(°6Ñ*:Ð":Ñ;Ñ<ð(#ØIQÐRZÑI[ð(#à	÷(#r*   r   c                   óÀ   ‡ — e Zd ZU dZdZeed<   eed<   dZdZe	e
   ed<   dZe	e
   ed<   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   df   de	e   d	efd„Zˆ xZS )ÚScaleInvariantSignalNoiseRatioa7  Calculate `Scale-invariant signal-to-noise ratio`_ (SI-SNR) metric for evaluating quality of audio.

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

    - ``preds`` (:class:`~torch.Tensor`): float tensor with shape ``(...,time)``
    - ``target`` (:class:`~torch.Tensor`): float tensor with shape ``(...,time)``

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

    - ``si_snr`` (:class:`~torch.Tensor`): float scalar tensor with average SI-SNR value over samples

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

    Raises:
        TypeError:
            if target and preds have a different shape

    Example:
        >>> import torch
        >>> from torch import tensor
        >>> from torchmetrics.audio import ScaleInvariantSignalNoiseRatio
        >>> target = tensor([3.0, -0.5, 2.0, 7.0])
        >>> preds = tensor([2.5, 0.0, 2.0, 8.0])
        >>> si_snr = ScaleInvariantSignalNoiseRatio()
        >>> si_snr(preds, target)
        tensor(15.0918)

    TÚ
sum_si_snrr   Nr   r   r   r   c                 ó˜   •— t        ‰| �  di |¤Ž | j                  dt        d«      d¬«       | j                  dt        d«      d¬«       y )NrK   r   r   r   r   r   r"   )r#   r$   r%   r   )r'   r   r(   s     €r)   r$   z'ScaleInvariantSignalNoiseRatio.__init__¸   sA   ø€ ô 	‰ÑÑ"˜6Ò"à�‰�|¬V°C«[ÈˆÔOØ�‰�w¬¨q«	À%ˆÕHr*   r+   r,   c                 óª   — t        ||¬«      }| xj                  |j                  «       z  c_        | xj                  |j	                  «       z  c_        y)r/   )r+   r,   N)r
   rK   r   r   r0   )r'   r+   r,   Úsi_snr_batchs       r)   r2   z%ScaleInvariantSignalNoiseRatio.updateÁ   s<   € ä9ÀÈfÔUˆà�Š˜<×+Ñ+Ó-Ñ-�Ø�
Š
�l×(Ñ(Ó*Ñ*Ž
r*   c                 ó4   — | j                   | j                  z  S r4   )rK   r   r5   s    r)   r6   z&ScaleInvariantSignalNoiseRatio.computeÈ   s   € à�‰ §¡Ñ+Ð+r*   r7   r8   c                 ó&   — | j                  ||«      S )a6  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.audio import ScaleInvariantSignalNoiseRatio
            >>> metric = ScaleInvariantSignalNoiseRatio()
            >>> metric.update(torch.rand(4), torch.rand(4))
            >>> fig_, ax_ = metric.plot()

        .. plot::
            :scale: 75

            >>> # Example plotting multiple values
            >>> import torch
            >>> from torchmetrics.audio import ScaleInvariantSignalNoiseRatio
            >>> metric = ScaleInvariantSignalNoiseRatio()
            >>> values = [ ]
            >>> for _ in range(10):
            ...     values.append(metric(torch.rand(4), torch.rand(4)))
            >>> fig_, ax_ = metric.plot(values)

        r:   r<   s      r)   r=   z#ScaleInvariantSignalNoiseRatio.plotÌ   ó   € ðL �z‰z˜#˜rÓ"Ð"r*   r?   )r@   rA   rB   rC   r   r   rE   r   r   r   rF   r   r   r$   r2   r6   r   r   r   r   r=   rG   rH   s   @r)   rJ   rJ   ’   s±   ø… ñð< ÐØÓØƒMØÐØ(,Ð�h˜u‘oÓ,Ø(,Ð�h˜u‘oÓ,ðIàðIð 
õIð+˜Fð +¨Fð +°tó +ð,˜ó ,ñ&#˜˜f h¨vÑ&6¸Ð<Ñ=ð &#È(ÐS[ÑJ\ð &#Ðhv÷ &#r*   rJ   c                   óÈ   ‡ — e Zd ZU dZdZeed<   eed<   dZdZe	e
   ed<   dZe	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   df   de	e   d
efd„Zˆ xZS )Ú%ComplexScaleInvariantSignalNoiseRatioaÖ  Calculate `Complex scale-invariant signal-to-noise ratio`_ (C-SI-SNR) metric for evaluating quality of audio.

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

    - ``preds`` (:class:`~torch.Tensor`): real float tensor with shape ``(...,frequency,time,2)`` or complex float
      tensor with shape ``(..., frequency,time)``

    - ``target`` (:class:`~torch.Tensor`): real float tensor with shape ``(...,frequency,time,2)`` or complex float
      tensor with shape ``(..., frequency,time)``

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

    - ``c_si_snr`` (:class:`~torch.Tensor`): float scalar tensor with average C-SI-SNR value over samples

    Args:
        zero_mean: if to zero mean target and preds or not
        kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info.

    Raises:
        ValueError:
            If ``zero_mean`` is not an bool
        TypeError:
            If ``preds`` is not the shape (..., frequency, time, 2) (after being converted to real if it is complex).
            If ``preds`` and ``target`` does not have the same shape.

    Example:
        >>> from torch import randn
        >>> from torchmetrics.audio import ComplexScaleInvariantSignalNoiseRatio
        >>> preds = randn((1,257,100,2))
        >>> target = randn((1,257,100,2))
        >>> c_si_snr = ComplexScaleInvariantSignalNoiseRatio()
        >>> c_si_snr(preds, target)
        tensor(-38.8832)

    TÚ
ci_snr_sumÚnumNr   r   r   r   r   c                 óâ   •— t        ‰| �  di |¤Ž t        |t        «      st	        d|› �«      ‚|| _        | j                  dt        d«      d¬«       | j                  dt        d«      d¬«       y )	Nz5Expected argument `zero_mean` to be an bool, but got rT   r   r   r   rU   r   r"   )r#   r$   Ú
isinstancerD   Ú
ValueErrorr   r%   r   r&   s      €r)   r$   z.ComplexScaleInvariantSignalNoiseRatio.__init__!  sg   ø€ ô
 	‰ÑÑ"˜6Ò"Ü˜)¤TÔ*ÜÐTÐU^ÐT_Ð`ÓaÐaØ"ˆŒà�‰�|¬V°C«[ÈˆÔOØ�‰�u¤f¨Q£iÀˆÕFr*   r+   r,   c                 óÀ   — t        ||| j                  ¬«      }| xj                  |j                  «       z  c_        | xj                  |j                  «       z  c_        yr.   )r	   r   rT   r   rU   r0   )r'   r+   r,   Úvs       r)   r2   z,ComplexScaleInvariantSignalNoiseRatio.update.  s?   € ä6¸UÈ6Ð]a×]kÑ]kÔlˆà�Š˜1Ÿ5™5›7Ñ"�Ø�Š�A—G‘G“IÑŽr*   c                 ó4   — | j                   | j                  z  S r4   )rT   rU   r5   s    r)   r6   z-ComplexScaleInvariantSignalNoiseRatio.compute5  s   € à�‰ §¡Ñ)Ð)r*   r7   r8   c                 ó&   — | j                  ||«      S )az  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.audio import ComplexScaleInvariantSignalNoiseRatio
            >>> metric = ComplexScaleInvariantSignalNoiseRatio()
            >>> metric.update(torch.rand(1,257,100,2), torch.rand(1,257,100,2))
            >>> fig_, ax_ = metric.plot()

        .. plot::
            :scale: 75

            >>> # Example plotting multiple values
            >>> import torch
            >>> from torchmetrics.audio import ComplexScaleInvariantSignalNoiseRatio
            >>> metric = ComplexScaleInvariantSignalNoiseRatio()
            >>> values = [ ]
            >>> for _ in range(10):
            ...     values.append(metric(torch.rand(1,257,100,2), torch.rand(1,257,100,2)))
            >>> fig_, ax_ = metric.plot(values)

        r:   r<   s      r)   r=   z*ComplexScaleInvariantSignalNoiseRatio.plot9  rQ   r*   r>   r?   )r@   rA   rB   rC   r   r   rE   r   r   r   rF   r   rD   r   r$   r2   r6   r   r   r   r   r=   rG   rH   s   @r)   rS   rS   õ   sÂ   ø… ñ"ðH ÐØÓØ	ƒKØÐØ(,Ð�h˜u‘oÓ,Ø(,Ð�h˜u‘oÓ,ð  ñGàðGð ðGð 
õ	Gð˜Fð ¨Fð °tó ð*˜ó *ñ&#˜˜f h¨vÑ&6¸Ð<Ñ=ð &#È(ÐS[ÑJ\ð &#Ðhv÷ &#r*   rS   N)Úcollections.abcr   Útypingr   r   r   Útorchr   r   Ú!torchmetrics.functional.audio.snrr	   r
   r   Útorchmetrics.metricr   Útorchmetrics.utilities.importsr   Útorchmetrics.utilities.plotr   r   Ú__doctest_skip__r   rJ   rS   r"   r*   r)   ú<module>re      s_   ðõ %ß 'Ñ 'ç  ÷ñ õ
 'Ý @ß @áòÐôk#�vô k#ô\`# Vô `#ôFj#¨Fõ j#r*   