Ë
    þÍ:jÒ  ã                   ó²   — 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 d dlmZ d dlmZmZmZ d dlmZmZ  eeeg«      sdd	gZnesd	gZ G d
„ de«      Zy)é    )ÚSequence)ÚAnyÚOptionalÚUnion)ÚTensorÚtensor)Ú_srmr_arg_validateÚ,speech_reverberation_modulation_energy_ratio)ÚMetric)Ú_GAMMATONE_AVAILABLEÚ_MATPLOTLIB_AVAILABLEÚ_TORCHAUDIO_AVAILABLE)Ú_AX_TYPEÚ_PLOT_OUT_TYPEÚ(SpeechReverberationModulationEnergyRatioz-SpeechReverberationModulationEnergyRatio.plotc                   ó  ‡ — e Zd ZU dZeed<   eed<   dZeed<   dZeed<   dZ	eed<   d	Z
ee   ed
<   d	Zee   ed<   	 	 	 	 	 	 ddededededee   dedededd	fˆ fd„Z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 )r   a;	  Calculate `Speech-to-Reverberation Modulation Energy Ratio`_ (SRMR).

    SRMR is a non-intrusive metric for speech quality and intelligibility based on
    a modulation spectral representation of the speech signal.
    This code is translated from SRMRToolbox and `SRMRpy`_.

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

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

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

    - ``srmr`` (:class:`~torch.Tensor`): float scaler tensor

    .. hint::
        Using this metrics requires you to have ``gammatone`` and ``torchaudio`` installed.
        Either install as ``pip install torchmetrics[audio]`` or ``pip install torchaudio``
        and ``pip install git+https://github.com/detly/gammatone``.

    .. attention::
        This implementation is experimental, and might not be consistent with the matlab
        implementation SRMRToolbox, especially the fast implementation.
        The slow versions, a) ``fast=False, norm=False, max_cf=128``, b) ``fast=False, norm=True, max_cf=30``,
        have a relatively small inconsistency.

    Args:
        fs: the sampling rate
        n_cochlear_filters: Number of filters in the acoustic filterbank
        low_freq: determines the frequency cutoff for the corresponding gammatone filterbank.
        min_cf: Center frequency in Hz of the first modulation filter.
        max_cf: Center frequency in Hz of the last modulation filter. If None is given,
            then 30 Hz will be used for `norm==False`, otherwise 128 Hz will be used.
        norm: Use modulation spectrum energy normalization
        fast: Use the faster version based on the gammatonegram.
            Note: this argument is inherited from `SRMRpy`_. As the translated code is based to pytorch,
            setting `fast=True` may slow down the speed for calculating this metric on GPU.

    Raises:
        ModuleNotFoundError:
            If ``gammatone`` or ``torchaudio`` package is not installed

    Example:
        >>> from torch import randn
        >>> from torchmetrics.audio import SpeechReverberationModulationEnergyRatio
        >>> preds = randn(8000)
        >>> srmr = SpeechReverberationModulationEnergyRatio(8000)
        >>> srmr(preds)
        tensor(0.3191)

    ÚmsumÚtotalFÚfull_state_updateTÚis_differentiableÚhigher_is_betterNÚplot_lower_boundÚplot_upper_boundÚfsÚn_cochlear_filtersÚlow_freqÚmin_cfÚmax_cfÚnormÚfastÚkwargsÚreturnc           	      óL  •— t        ‰	| �  d	i |¤Ž t        rt        st	        d«      ‚t        |||||||¬«       || _        || _        || _        || _	        || _
        || _        || _        | j                  dt        d«      d¬«       | j                  dt        d«      d¬«       y )
Na  speech_reverberation_modulation_energy_ratio requires you to have `gammatone` and `torchaudio>=0.10` installed. Either install as ``pip install torchmetrics[audio]`` or ``pip install torchaudio>=0.10`` and ``pip install git+https://github.com/detly/gammatone``)r   r   r   r   r   r   r    r   g        Úsum)ÚdefaultÚdist_reduce_fxr   r   © )ÚsuperÚ__init__r   r   ÚModuleNotFoundErrorr	   r   r   r   r   r   r   r    Ú	add_stater   )
Úselfr   r   r   r   r   r   r    r!   Ú	__class__s
            €úl/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchmetrics/audio/srmr.pyr)   z1SpeechReverberationModulationEnergyRatio.__init__a   s­   ø€ ô 	‰ÑÑ"˜6Ò"Ý$Õ,@Ü%ðnóð ô
 	ØØ1ØØØØØõ	
ð ˆŒØ"4ˆÔØ ˆŒØˆŒØˆŒØˆŒ	ØˆŒ	à�‰�v¤v¨c£{À5ˆÔIØ�‰�w¬¨q«	À%ˆÕHó    Úpredsc           
      ó†  — t        || j                  | j                  | j                  | j                  | j
                  | j                  | j                  «      j                  | j                  j                  «      }| xj                  |j                  «       z  c_	        | xj                  |j                  «       z  c_        y)zUpdate state with predictions.N)r
   r   r   r   r   r   r   r    Útor   Údevicer$   r   Únumel)r,   r0   Úmetric_val_batchs      r.   Úupdatez/SpeechReverberationModulationEnergyRatio.updateˆ   s�   € äGØ�4—7‘7˜D×3Ñ3°T·]±]ÀDÇKÁKÐQU×Q\ÑQ\Ð^b×^gÑ^gÐim×irÑiró
ç
‰"ˆT�Y‰Y×ÑÓ
ð 	ð 	�	Š	Ð%×)Ñ)Ó+Ñ+�	Ø�
Š
Ð&×,Ñ,Ó.Ñ.Ž
r/   c                 ó4   — | j                   | j                  z  S )zCompute metric.)r   r   )r,   s    r.   Úcomputez0SpeechReverberationModulationEnergyRatio.compute‘   s   € à�y‰y˜4Ÿ:™:Ñ%Ð%r/   ÚvalÚaxc                 ó&   — | j                  ||«      S )aN  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 SpeechReverberationModulationEnergyRatio
            >>> metric = SpeechReverberationModulationEnergyRatio(8000)
            >>> metric.update(torch.rand(8000))
            >>> fig_, ax_ = metric.plot()

        .. plot::
            :scale: 75

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

        )Ú_plot)r,   r9   r:   s      r.   Úplotz-SpeechReverberationModulationEnergyRatio.plot•   s   € ðL �z‰z˜#˜rÓ"Ð"r/   )é   é}   é   NFF)NN)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   Ú__annotations__r   Úboolr   r   r   r   Úfloatr   Úintr   r)   r6   r8   r   r   r   r   r=   Ú__classcell__)r-   s   @r.   r   r   %   s%  ø… ñ1ðf ƒLØƒMØ#Ð�tÓ#Ø"Ð�tÓ"Ø!Ð�dÓ!Ø(,Ð�h˜u‘oÓ,Ø(,Ð�h˜u‘oÓ,ð
 #%ØØØ"&ØØñ%Iàð%Ið  ð%Ið ð	%Ið
 ð%Ið ˜‘ð%Ið ð%Ið ð%Ið ð%Ið 
õ%IðN/˜Fð / tó /ð&˜ó &ñ&#˜˜f h¨vÑ&6¸Ð<Ñ=ð &#È(ÐS[ÑJ\ð &#Ðhv÷ &#r/   N)Úcollections.abcr   Útypingr   r   r   Útorchr   r   Ú"torchmetrics.functional.audio.srmrr	   r
   Útorchmetrics.metricr   Útorchmetrics.utilities.importsr   r   r   Útorchmetrics.utilities.plotr   r   ÚallÚ__doctest_skip__r   r'   r/   r.   ú<module>rS      sb   ðõ %ß 'Ñ 'ç  ÷õ '÷ñ ÷
 Aá
Ð Ð"7Ð8Ô9ØBÐDsÐtÑÙ	ØGÐHÐôV#¨võ V#r/   