Ë
    þÍ:jZ7  ã                   ó>  — d dl mZ d dlmZmZ d dl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 eresd	gZ ed
¬«      dedededej&                  def
d„«       Z ed
¬«      dedededej&                  def
d„«       Z ed
¬«      dedededededej&                  deeeeef   fd„«       Zd,dedee   defd„Zdededefd„Zd-ded edefd!„Zd"ed#ed$edefd%„Z	 	 	 	 	 	 d.d&eded'edededee   d(ed)edefd*„Z	 	 	 	 	 	 d/ded'edededee   d(ed)eddfd+„Zy)0é    )Ú	lru_cache)ÚceilÚpi)ÚOptionalN)ÚTensor)Úpad)Úrank_zero_warn)Ú_GAMMATONE_AVAILABLEÚ_TORCHAUDIO_AVAILABLEÚ,speech_reverberation_modulation_energy_ratioéd   )ÚmaxsizeÚlow_freqÚfsÚ	n_filtersÚdeviceÚreturnc                 ó€   — ddl m} d}d}d} |||| «      |z  |z  ||z  z   d|z  z  }t        j                  ||¬«      S )Nr   )Úcentre_freqsgó<¸;k‡"@g33333³8@é   ©r   )Úgammatone.filtersr   ÚtorchÚtensor)	r   r   r   r   r   Úear_qÚmin_bwÚorderÚerbss	            úw/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchmetrics/functional/audio/srmr.pyÚ
_calc_erbsr    $   sR   € å.à€EØ€FØ€EÙ˜"˜i¨Ó2°UÑ:¸uÑDÀvÈuÁ}ÑTÐZ[Ð^cÑZcÑd€DÜ�<‰<˜ VÔ,Ð,ó    Ú	num_freqsÚcutoffc                 óf   — ddl m}m}  || ||«      } || |«      }t        j                  ||¬«      S )Nr   )r   Úmake_erb_filtersr   )r   r   r%   r   r   )r   r"   r#   r   r   r%   ÚcfsÚfcoefss           r   Ú_make_erb_filtersr(   /   s0   € ç@á
�r˜9 fÓ
-€CÙ˜b #Ó&€FÜ�<‰<˜ vÔ.Ð.r!   Úmin_cfÚmax_cfÚnÚqc           
      ó  — || z  d|dz
  z  z  }t        j                  |t         j                  ¬«      }| |d<   t        d|«      D ]  }||dz
     |z  ||<   Œ dt        dt
        dt        fd„}	t        j                  d	t        z  |z  |z  D �
cg c]  }
 |	|
|«      ‘Œ c}
d¬
«      }dt        dt        dt
        dt        t        t        f   fd„}|j                  |¬«      }|j                  |¬«      } ||||«      \  }}||||fS c c}
w )Nç      ð?r   ©Údtyper   Úw0r,   r   c                 óF  — t        j                  | dz  «      } | |z  }t        j                  |d| gt         j                  ¬«      }t        j                  d|z   | dz  z   d| dz  z  dz
  d|z
  | dz  z   gt         j                  ¬«      }t        j                  ||gd¬«      S )Né   r   r/   r   ©Údim)r   Útanr   Úfloat64Ústack)r1   r,   Úb0ÚbÚas        r   Ú_make_modulation_filterzK_compute_modulation_filterbank_and_cutoffs.<locals>._make_modulation_filterC   s�   € Ü�Y‰Y�r˜A‘vÓˆØ�!‰VˆÜ�L‰L˜"˜a " ˜¬U¯]©]Ô;ˆÜ�L‰L˜1˜r™6 B¨¡E™>¨Q°°Q±©Y¸©]¸aÀ"¹fÀrÈ1Áu¹nÐNÔV[×VcÑVcÔdˆÜ�{‰{˜A˜q˜6 qÔ)Ð)r!   r3   r4   r&   r   c                 ó¦   — dt         z  | z  |z  }t        j                  |dz  «      |z  }| ||z  dt         z  z  z
  }| ||z  dt         z  z  z   }||fS )Nr3   )r   r   r6   )r&   r   r,   r1   r9   ÚllÚrrs          r   Ú_calc_cutoffszA_compute_modulation_filterbank_and_cutoffs.<locals>._calc_cutoffsL   sb   € à”‰V�c‰\˜BÑˆÜ�Y‰Y�r˜A‘vÓ Ñ"ˆØ�B˜‘G˜q¤2™vÑ&Ñ'ˆØ�B˜‘G˜q¤2™vÑ&Ñ'ˆØ�2ˆvˆr!   r   )r   Úzerosr7   Úranger   Úintr8   r   ÚfloatÚtupleÚto)r)   r*   r+   r   r,   r   Úspacing_factorr&   Úkr<   r1   Úmfbr@   r>   r?   s                  r   Ú*_compute_modulation_filterbank_and_cutoffsrJ   8   s"  € ð
 ˜v‘o¨3°!°a±%©=Ñ9€NÜ
�+‰+�aœuŸ}™}Ô
-€CØ€Cˆ�FÜ�1�a‹[ò -ˆØ�Q˜‘U‘˜nÑ,ˆˆAŠð-ð*¤Fð *¬sð *´vó *ô �+‰+ÀÄBÁÈÁÈrÑ@QÖR¸"Ñ.¨r°1Õ5ÒRÐXYÔ
Z€Cðœ6ð ¤uð ´ð ¼¼vÄv¸~Ñ9Nó ð �&‰&˜ˆ&Ó
€CØ
�&‰&˜ˆ&Ó
€CÙ˜3  AÓ&�F€BˆØ��R˜ÐÐùò Ss   ÂDÚxc                 ó  — | j                  «       rt        d«      ‚|€%| j                  d   }|dz  rt        |dz  «      dz  }|dk  rt        d«      ‚t        j
                  j                  | |d¬«      }t	        j                  || j                  | j                  d¬«      }|d	z  dk(  rd
x|d<   ||d	z  <   d	|d
|d	z   nd
|d<   d	|d
|d
z   d	z   t        j
                  j                  ||z  d¬«      }|dd | j                  d   …f   S )Nzx must be real.éÿÿÿÿé   r   zN must be positive.)r+   r5   F)r0   r   Úrequires_gradr3   r   r4   .)
Ú
is_complexÚ
ValueErrorÚshaper   r   ÚfftrA   r0   r   Úifft)rK   r+   Úx_fftÚhÚys        r   Ú_hilbertrX   Z   s  € Ø‡|�|„~ÜÐ*Ó+Ð+Ø€yØ�G‰G�B‰KˆàˆrŠ6Ü�Q˜‘V“˜rÑ!ˆAØˆA‚vÜÐ.Ó/Ð/ä�I‰I�M‰M˜!˜q bˆMÓ)€EÜ�‰�A˜QŸW™W¨Q¯X©XÀUÔK€Aàˆ1�u�‚zØÐˆˆ!‰ˆq��a‘‰yØˆˆ!ˆa�1‰f‰àˆˆ!‰Øˆˆ!ˆq�1‰u˜‰lÐä�	‰	�‰�u˜q‘y bˆÓ)€AØˆS�-�A—G‘G˜B‘K�-ÐÑ Ð r!   ÚwaveÚcoefsc                 óÂ  — ddl m} | j                  \  }}| j                  |j                  ¬«      j                  |d|«      } | j                  d|j                  d   d«      } |dd…df   }|dd…df   }|dd…d	f   }|dd…d
f   }|dd…df   }	|dd…dd…f   }
 || |
|d¬«      } |||
|d¬«      } |||
|d¬«      } |||
|	d¬«      }||j                  ddd«      z  S )zŸTranslated from gammatone package.

    Args:
        wave: shape [B, time]
        coefs: shape [N, 10]

    Returns:
        Tensor: shape [B, N, time]

    r   ©Úlfilterr/   r   rM   Né	   )r   r   é   )r   r3   r_   )r   é   r_   )r   é   r_   é   T)Úbatching)Útorchaudio.functional.filteringr]   rR   rF   r0   ÚreshapeÚexpand)rY   rZ   r]   Ú	num_batchÚtimeÚgainÚas1Úas2Úas3Úas4ÚbsÚy1Úy2Úy3Úy4s                  r   Ú_erb_filterbankrs   s   sÿ   € õ 8à—j‘j�O€IˆtØ�7‰7˜Ÿ™ˆ7Ó%×-Ñ-¨i¸¸DÓA€DØ�;‰;�r˜5Ÿ;™; q™>¨2Ó.€Dà’�A�‰;€DØ
’�9�Ñ
€CØ
’�9�Ñ
€CØ
’�9�Ñ
€CØ
’�9�Ñ
€CØ	Šq�!�A�#ˆv‰€Bá	��r˜3¨Ô	.€BÙ	��R˜ tÔ	,€BÙ	��R˜ tÔ	,€BÙ	��R˜ tÔ	,€BØ�—‘˜Q  AÓ&Ñ&Ð&r!   ÚenergyÚdrangec                 ó"  — t        j                  | dd¬«      j                  dd¬«      j                  }|j                  dd¬«      j                  }|d| dz  z  z  }t        j                  | |k  || «      } t        j                  | |kD  || «      S )z”Normalize energy to a dynamic range of 30 dB.

    Args:
        energy: shape [B, N_filters, 8, n_frames]
        drange: dynamic range in dB

    r   T©r5   Úkeepdimr3   r`   g      $@)r   ÚmeanÚmaxÚvaluesÚwhere)rt   ru   Úpeak_energyÚ
min_energys       r   Ú_normalize_energyr   ’   sˆ   € ô —*‘*˜V¨°DÔ9×=Ñ=À!ÈTÐ=ÓR×YÑY€KØ—/‘/ a°�/Ó6×=Ñ=€KØ˜t¨¨°$©Ñ7Ñ7€JÜ�[‰[˜ *Ñ,¨j¸&ÓA€FÜ�;‰;�v Ñ+¨[¸&ÓAÐAr!   ÚbwÚ
avg_energyÚcutoffsc                 ó  — |d   | k  r|d   | kD  rd}n<|d   | k  r|d   | kD  rd}n)|d   | k  r|d   | kD  rd}n|d   | k  rd}nt        d«      ‚t        j                  |dd…dd…f   «      t        j                  |dd…d|…f   «      z  S )zCalculate srmr score.ra   r_   rb   é   é   z7Something wrong with the cutoffs compared to bw values.N)rQ   r   Úsum)r€   r�   r‚   Úkstars       r   Ú_cal_srmr_scorerˆ   ¡   s§   € à�‰
�bÒ˜w q™z¨BšØ‰Ø
�!‰*˜Ò
 ¨¡¨b¢Ø‰Ø
�!‰*˜Ò
 ¨¡¨b¢Ø‰Ø	�‰�rÒ	Ø‰äÐRÓSÐSÜ�9‰9�Z¢ 2 A 2 Ñ&Ó'¬%¯)©)°JºqÀ!ÀEÀ'¸zÑ4JÓ*KÑKÐKr!   ÚpredsÚn_cochlear_filtersÚnormÚfastc           	      ó 
  — t         rt        st        d«      ‚ddlm} ddlm}	 t        |||||||¬«       | j                  }
t        |
«      dk(  r| j                  dd«      n| j                  d|
d   «      } | j                  \  }}t        j                  | «      sI| j                  t        j                  «      t        j                  | j                   «      j"                  z  } | j%                  «       j#                  dd¬	«      j&                  }t        j(                  |dkD  |t        j*                  d
|j                   |j,                  ¬«      «      }| |z  } d}d}|r±t/        d«       d}g }| j1                  «       j3                  «       j5                  «       }t7        |«      D ]6  } |||   |dd||«      }|j9                  t        j*                  |«      «       Œ8 t        j:                  |d¬«      j                  | j,                  ¬«      }nCt=        |||| j,                  ¬«      }t        j$                  t?        tA        | |«      «      «      }|}tC        ||z  «      }tC        ||z  «      }|€|rdnd}tE        ||d|d| j,                  ¬«      \  }}}}tG        d||z
  |z  z   «      }t        jH                  |dz   t        j                  | j,                  ¬«      dd } |	|jK                  d«      jM                  dd|j                  d   d«      |dd…ddd…f   |dd…ddd…f   dd¬«      }dt#        tC        ||z  «      |z  |z
  ||z
  «      f} tO        || dd¬«      }!|!jQ                  d||«      }"|"dd|…dd…f   |z  dz  jS                  d¬«      }#|rtU        |#«      }#t        jV                  tY        |||| j,                  ¬«      «      }$t        jZ                  |#d¬«      }%t        jR                  |%j                  |d«      d¬«      }&t        jR                  |%d¬«      }'|'d z  |&j                  dd«      z  }(|(j]                  d«      j_                  d«      })t        j`                  |)d!kD  j_                  d«      dk(  «      dd…df   }*|$|*   }+g }t7        |«      D ]'  }tc        |+|   |%|   |¬"«      },|j9                  |,«       Œ) t        j:                  |«      },t        |
«      dkD  r |,j                  |
dd Ž S |,S )#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`_.

    Args:
        preds: shape ``(..., time)``
        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.

    .. hint::
        Usingsing 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.

    Returns:
        Scalar tensor with srmr value with shape ``(...)``

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

    Example:
        >>> from torch import randn
        >>> from torchmetrics.functional.audio import speech_reverberation_modulation_energy_ratio
        >>> preds = randn(8000)
        >>> speech_reverberation_modulation_energy_ratio(preds, 8000)
        tensor([0.3191], dtype=torch.float64)

    a  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   )Ú
fft_gtgramr\   ©r   rŠ   r   r)   r*   r‹   rŒ   r   rM   Trw   r.   )r0   r   gü©ñÒMbÐ?gü©ñÒMb°?z:`fast=True` may slow down the speed of SRMR metric on GPU.g      y@g{®Gáz„?g{®Gázd?r4   r   Né   é€   r…   r3   )r+   r   r,   r   éþÿÿÿF)Úclamprc   Úconstant)r   ÚmodeÚvalue.r   éZ   )r‚   )2r   r
   ÚModuleNotFoundErrorÚgammatone.fftweightrŽ   rd   r]   Ú_srmr_arg_validaterR   Úlenre   r   Úis_floating_pointrF   r7   Úfinfor0   rz   Úabsr{   r|   r   r   r	   ÚdetachÚcpuÚnumpyrB   Úappendr8   r(   rX   rs   r   rJ   rC   Úhamming_windowÚ	unsqueezerf   r   Úunfoldr†   r   Úflipudr    ry   ÚflipÚcumsumÚnonzerorˆ   )-r‰   r   rŠ   r   r)   r*   r‹   rŒ   rŽ   r]   rR   rg   rh   Úmax_valsÚval_normÚ
w_length_sÚw_inc_sÚmfsÚtempÚpreds_npr:   Úgt_env_bÚgt_envr'   Úw_lengthÚw_incÚ_Úmfr‚   Ú
num_framesÚwÚmod_outÚpaddingÚmod_out_padÚmod_out_framert   r   r�   Útotal_energyÚ	ac_energyÚac_percÚac_perc_cumsumÚk90perc_idxr€   Úscores-                                                r   r   r   °   s–  € õn !Õ(<Ü!ðjó
ð 	
õ
 /Ý7äØØ-ØØØØØõð �K‰K€EÜ$'¨£J°!¢OˆE�M‰M˜!˜RÔ ¸¿¹ÀrÈ5ÐQSÉ9Ó9U€EØ—k‘k�O€Iˆtä×"Ñ" 5Ô)Ø—‘œŸ™Ó'¬%¯+©+°e·k±kÓ*B×*FÑ*FÑFˆð �y‰y‹{�‰ 2¨tˆÓ4×;Ñ;€HÜ�{‰{Ø�1‰ØÜ�‰�S §¡°x·±ÔGó€Hð
 �HÑ€Eà€JØ€GáÜÐSÔTØˆØˆØ—<‘<“>×%Ñ%Ó'×-Ñ-Ó/ˆÜ�yÓ!ò 	0ˆAÙ! (¨1¡+¨r°5¸&ÐBTÐV^Ó_ˆHØ�K‰KœŸ™ XÓ.Õ/ð	0ô —‘˜T qÔ)×,Ñ,°E·L±LÐ,ÓA‰ä" 2Ð'9¸8ÈEÏLÉLÔYˆÜ—‘œ8¤O°E¸6Ó$BÓCÓDˆØˆä�J Ñ$Ó%€HÜ�˜3‘Ó€Eð €~Ù‘ ˆÜBØ�˜!  q°·±ôÑ€A€rˆ7�Aô �Q˜$ ™/¨eÑ3Ñ3Ó4€JÜ×Ñ˜X¨™\´·±ÀuÇ|Á|ÔTÐUXÐVXÐY€AÙØ×Ñ˜Ó×#Ñ# B¨¨B¯H©H°Q©K¸Ó<¸bÂÀAÂqÀ¹kÈ2ÊaÐQRÒTUÈgÉ;Ð^cÐnrô€Gð ”#”d˜4 %™<Ó(¨5Ñ0°4Ñ7¸ÀD¹ÓIÐJ€GÜ�g 7°À1ÔE€KØ×&Ñ& r¨8°UÓ;€MØ˜S + : +ªqÐ0Ñ1°AÑ5¸!Ñ;×@Ñ@ÀRÐ@ÓH€FáÜ" 6Ó*ˆä�<‰<œ
 8¨RÐ1CÈEÏLÉLÔYÓZ€Dä—‘˜F¨Ô+€JÜ—9‘9˜Z×/Ñ/°	¸2Ó>ÀBÔG€LÜ—	‘	˜*¨!Ô,€IØ˜#‰o × 4Ñ 4°R¸Ó ;Ñ;€GØ—\‘\ "Ó%×,Ñ,¨RÓ0€NÜ—-‘- °"Ñ!4× <Ñ <¸RÓ @ÀAÑ EÓFÂqÈ!ÀtÑL€KØ	ˆkÑ	€Bà€DÜ�9Óò ˆÜ  1¡ z°!¡}¸gÔFˆØ�‰�EÕðô �K‰K˜Ó€Eä),¨U«°aªˆ=ˆ5�=‰=˜%  ˜*Ð%ÐB¸UÐBr!   c                 óö  — t        | t        «      r| dkD  st        d| › �«      ‚t        |t        «      r|dkD  st        d|› �«      ‚t        |t        t        f«      r|dkD  st        d|› �«      ‚t        |t        t        f«      r|dkD  st        d|› �«      ‚|�)t        |t        t        f«      r|dkD  st        d|› �«      ‚t        |t        «      st        d«      ‚t        |t        «      st        d	«      ‚y)
a9  Validate the arguments for speech_reverberation_modulation_energy_ratio.

    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,
        norm: Use modulation spectrum energy normalization
        fast: Use the faster version based on the gammatonegram.

    r   z;Expected argument `fs` to be an int larger than 0, but got zKExpected argument `n_cochlear_filters` to be an int larger than 0, but got zBExpected argument `low_freq` to be a float larger than 0, but got z@Expected argument `min_cf` to be a float larger than 0, but got Nz@Expected argument `max_cf` to be a float larger than 0, but got z+Expected argument `norm` to be a bool valuez+Expected argument `fast` to be a bool value)Ú
isinstancerC   rQ   rD   Úboolr�   s          r   rš   rš   E  s	  € ô* �rœ3Ô B¨¢FÜÐVÐWYÐVZÐ[Ó\Ð\ÜÐ)¬3Ô/Ð4FÈÒ4JÜØYÐZlÐYmÐnó
ð 	
ô ˜¤5¬# ,Ô/°XÀ²\ÜÐ]Ð^fÐ]gÐhÓiÐiÜ˜¤¬ Ô-°6¸A²:ÜÐ[Ð\bÐ[cÐdÓeÐeØÐ¤J¨v¼¼s°|Ô$DÈ&ÐSTÊ*ÜÐ[Ð\bÐ[cÐdÓeÐeÜ�dœDÔ!ÜÐFÓGÐGÜ�dœDÔ!ÜÐFÓGÐGð "r!   )N)g      >@)é   é}   ra   NFF)rÆ   rÇ   ra   r‘   FF)Ú	functoolsr   Úmathr   r   Útypingr   r   r   Útorch.nn.functionalr   Útorchmetrics.utilitiesr	   Útorchmetrics.utilities.importsr
   r   Ú__doctest_skip__rD   rC   r   r    r(   rE   rJ   rX   rs   r   rˆ   rÅ   r   rš   © r!   r   ú<module>rÐ      sŽ  ðõ$  ß Ý ã Ý Ý #å 1÷ñ
 Ñ$8ØFÐGÐñ �3Ôð-˜ð - Cð -°Cð -ÀÇÁð -ÐRXò -ó ð-ñ �3Ôð/˜#ð /¨#ð /°uð /ÀeÇlÁlð /ÐW]ò /ó ð/ñ �3ÔðØðØ ðØ%(ðØ.3ðØ8;ðØEJÇ\Á\ðà
ˆ6�6˜6 6Ð)Ñ*òó ðñB!�ð !˜8 C™=ð !°Fó !ð2'˜&ð '¨ð '°Fó 'ñ>B˜fð B¨eð B¸vó BðL˜ð L¨Fð L¸Vð LÈó Lð$ !ØØØ"ØØñRCØðRCàðRCð ðRCð ð	RCð
 ðRCð �U‰OðRCð ðRCð ðRCð óRCðn !ØØØ!ØØñ$HØð$Hàð$Hð ð$Hð ð	$Hð
 �U‰Oð$Hð ð$Hð ð$Hð 
ô$Hr!   