a
    d8                     @   s  d dl mZ d dlmZmZ d dl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mZ ererd d	lmZ n
dZd
gZerd dlmZ d dlmZmZ nd\ZZZd
gZeddeeeeje	dddZeddeeeeje	dddZeddeeeeeejee	e	e	e	f dddZd0e	ee e	dddZ e	e	e	dddZ!d1e	ee	d d!d"Z"e	e	e	e	d#d$d%Z#d2e	eeeeee e$e$e	d*	d+d
Z%d3eeeeee e$e$dd-d.d/Z&dS )4    )	lru_cache)ceilpi)OptionalTupleN)Tensor)pad)rank_zero_warn)_GAMMATONE_AVAILABEL_TORCHAUDIO_AVAILABEL_TORCHAUDIO_GREATER_EQUAL_0_10)lfilter,speech_reverberation_modulation_energy_ratio)
fft_gtgram)centre_freqsmake_erb_filters)NNNd   )maxsize)low_freqfs	n_filtersdevicereturnc                 C   s>   d}d}d}t ||| | | ||  d|  }tj||dS )Ng<;k"@g333338@   r   )r   torchtensor)r   r   r   r   Zear_qZmin_bwordererbs r   k/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/torchmetrics/functional/audio/srmr.py
_calc_erbs/   s
    $r!   )r   	num_freqscutoffr   r   c                 C   s$   t | ||}t| |}tj||dS )Nr   )r   r   r   r   )r   r"   r#   r   cfsfcoefsr   r   r    _make_erb_filters8   s    
r&   )min_cfmax_cfnr   qr   r   c                    s   ||  d|d   }t j|t jd}| |d< td|D ]}||d  | ||< q6tttddd t j fdd	d
t | | D dd}	tttt	ttf ddd}
|j
|d}|	j
|d}	|
||\}}||	||fS )N      ?r   dtyper   )w0r*   r   c                 S   sz   t | d } | | }t j|d| gt jd}t jd| | d  d| d  d d| | d  gt jd}t j||gddS )N   r   r,   r   dim)r   tanr   float64stack)r.   r*   b0bar   r   r    _make_modulation_filterJ   s
    :zK_compute_modulation_filterbank_and_cutoffs.<locals>._make_modulation_filterc                    s   g | ]} |qS r   r   ).0r.   r8   r*   r   r    
<listcomp>Q       z>_compute_modulation_filterbank_and_cutoffs.<locals>.<listcomp>r/   r0   )r$   r   r*   r   c                 S   sR   dt  |  | }t|d | }| || dt    }| || dt    }||fS )Nr/   )r   r   r2   )r$   r   r*   r.   r5   llrrr   r   r    _calc_cutoffsS   s
    zA_compute_modulation_filterbank_and_cutoffs.<locals>._calc_cutoffsr   )r   zerosr3   ranger   intr4   r   floatr   to)r'   r(   r)   r   r*   r   Zspacing_factorr$   kZmfbr?   r=   r>   r   r:   r    *_compute_modulation_filterbank_and_cutoffs?   s    *rF   )xr)   r   c                 C   s   |   rtd|d u r:| jd }|d r:t|d d }|dkrJtdtjj| |dd}tj|| j| jdd}|d	 dkrd
 |d< ||d	 < d	|d
|d	 < nd
|d< d	|d
|d
 d	 < tjj	|| dd}|dd | jd f S )Nzx must be real.   r   zN must be positive.)r)   r1   F)r-   r   Zrequires_gradr/   r   r0   .)
Z
is_complex
ValueErrorshaper   r   Zfftr@   r-   r   Zifft)rG   r)   Zx_ffthyr   r   r    _hilberta   s"    
rN   )wavecoefsr   c                 C   s   | j \}}| j|jd|d|} | 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 }	t| |	|dd}
t|
|	|dd}t||	|dd}t||	|dd}||ddd S )zTranslated from gammatone package.

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

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

    r,   r   rH   r   N	   )r   r      )r   r/   rR   )r      rR   )r      rR      T)batching)rK   rD   r-   reshapeexpandr   )rO   rP   n_batchtimeZgainZas1Zas2Zas3Zas4bsy1y2Zy3Zy4r   r   r    _erb_filterbankz   s    
r^         >@)energydranger   c                 C   sb   t j| dddjdddj}|jdddj}|d| d   }t | |k || } t | |k|| S )zNormalize energy to a dynamic range of 30 dB.

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

    r   Tr1   Zkeepdimr/   rS   g      $@)r   meanmaxvalueswhere)r`   ra   Zpeak_energyZ
min_energyr   r   r    _normalize_energy   s
    rg   )bw
avg_energycutoffsr   c                 C   s   |d | kr|d | krd}nV|d | kr<|d | kr<d}n8|d | krZ|d | krZd}n|d | krld}nt dt|ddddf t|ddd|f  S )zCalculate srmr score.rT   rR   rU         z7Something wrong with the cutoffs compared to bw values.N)rJ   r   sum)rh   ri   rj   Zkstarr   r   r    _cal_srmr_score   s    rn      }   rT   F)	predsr   n_cochlear_filtersr   r'   r(   normfastr   c           +   	   C   s  t rtrtstdt|||||||d | j}t|dkrH| ddn| d|d } | j\}	}
t	| s| 
tjt| jj } |  jdddj}t|dk|tjd|j|jd}| | } d	}d
}|r@td d}g }|    }t|	D ]*}t|| |dd||}|t| qtj|ddj
| jd}n*t|||| jd}ttt| |}|}t|| }t|| }|du r|rdnd}t ||d|d| jd\}}}}t!d|
| |  }tj"|d tj| jddd }t#|$d%dd|jd d|dddddf |dddddf ddd}dtt|
| | |
 ||
 f}t&||ddd}|'d||} | dd|ddf | d j(dd}!|rt)|!}!t*t+|||| jd}"tj,|!dd}#tj(|#|	ddd}$tj(|#dd}%|%d |$dd }&|&-d.d}'t/|'dk.ddkdddf }(|"|( })g }t|	D ]&}t0|)| |#| |d }*||* qTt|}*t|dk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.

    .. note:: 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``.

    .. note::
        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 inconsistence.

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

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

    Example:
        >>> import torch
        >>> from torchmetrics.functional.audio import speech_reverberation_modulation_energy_ratio
        >>> g = torch.manual_seed(1)
        >>> preds = torch.randn(8000)
        >>> speech_reverberation_modulation_energy_ratio(preds, 8000)
        tensor([0.3354], 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   rr   r   r'   r(   rs   rt   r   rH   Trb   r+   )r-   r   gMb?gMb?z:`fast=True` may slow down the speed of SRMR metric on GPU.g      y@g{Gz?g{Gzd?r   r0   r   N      rl   r/   )r)   r   r*   r   F)clamprV   Zconstant)r   modevalue.r   Z   )rj   )1r   r   r
   ModuleNotFoundError_srmr_arg_validaterK   lenrW   r   Zis_floating_pointrD   r3   Zfinfor-   rd   absre   rf   r   r   r	   detachcpunumpyrA   r   appendr4   r&   rN   r^   r   rF   rB   Zhamming_windowr   Z	unsqueezerX   r   Zunfoldrm   rg   Zflipudr!   rc   ZflipZcumsumZnonzerorn   )+rq   r   rr   r   r'   r(   rs   rt   rK   rY   rZ   Zmax_valsZval_normZ
w_length_sZw_inc_sZmfstempZpreds_npr6   Zgt_env_bZgt_envr%   Zw_lengthZw_inc_Zmfrj   Zn_frameswZmod_outpaddingZmod_out_padZmod_out_framer`   r   ri   Ztotal_energyZ	ac_energyZac_percZac_perc_cumsumZk90perc_idxrh   Zscorer   r   r    r      s    7	(


 F"&$
rw   )r   rr   r   r'   r(   rs   rt   r   c                 C   s   t | tr| dks td|  t |tr2|dks@td| t |ttfrV|dksdtd| t |ttfrz|dkstd| |durt |ttfr|dkstd| t |tstdt |tstd	dS )
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)
isinstancerB   rJ   rC   boolru   r   r   r    r~   G  s     

r~   )N)r_   )ro   rp   rT   NFF)ro   rp   rT   rw   FF)'	functoolsr   mathr   r   typingr   r   r   r   Ztorch.nn.functionalr   Ztorchmetrics.utilitiesr	   Ztorchmetrics.utilities.importsr
   r   r   Ztorchaudio.functional.filteringr   Z__doctest_skip__Zgammatone.fftweightr   Zgammatone.filtersr   r   rC   rB   r   r!   r&   rF   rN   r^   rg   rn   r   r   r~   r   r   r   r    <module>   s|   
!             