
    iZ7                     >   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_ratiod   )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"@g333338@   r   )gammatone.filtersr   torchtensor)	r   r   r   r   r   ear_qmin_bwordererbss	            w/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/torchmetrics/functional/audio/srmr.py
_calc_erbsr    $   sR    .EFE"i2U:uDvu}TZ[^cZcdD<<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    @
r9f
-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    YYrAv!VLL"a"U]];LL1r6BE>QQY]a"fr1unNV[VcVcd{{Aq6q))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    Vc\BYYrAv"BGq2v&'BGq2v&'2v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#   
 vo3!a%=9N
++au}}
-CCF1a[QUn,A *F *s *v * ++Br@QR@Q".r15@QRXY
ZC6 u  vv~9N  &&&
C
&&&
C3A&FB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GGBKr6QVr!AAv.//IIMM!qbM)EAQWWQXXUKA1uz!qay!a1f!!q1ul		uqyb)AS-AGGBK-  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 )zTranslated 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jjOIt777%--iDAD;;r5;;q>2.DA;D
9
C
9
C
9
C
9
C	q!A#vB	r3	.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 )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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YK//a/6==Kt$77J[[*,j&AF;;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wqzB
!*
b
!*
b	r	RSS99Z2A2&'%))Jq!E'z4J*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}|rt/        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   gMb?gMb?z:`fast=True` may slow down the speed of SRMR metric on GPU.g      y@g{Gz?g{Gzd?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- KKE$'J!OEMM!R r5QS99UEkkOIt""5)'%++ekk*B*F*FF yy{2t4;;H{{1SxGH
 HEJGST<<>%%'--/y!A!(1+r5&BTV^_HKKX./ " Tq),,ELL,A"2'98ELLY8OE6$BCDJ$%H3E ~B!qAr7A Q$/e334JX\u||TUXVXYA##BBHHQK<bAqk2aQRTUg;^cnrG #d4%<(5047DIJGg71EK&&r8U;MS+:+q01A5!;@@R@HF"6*<<
8R1CELLYZDF+J99Z//	2>BGL		*!,I#o 4 4R ;;G\\"%,,R0N--"!4 < <R @A EFq!tLK	k	BD91z!}gFE  KKE),Ua=5==%*%BU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	   * r3BFVWYVZ[\\)3/4F4JYZlYmn
 	
 5#,/X\]^f]ghii-6A:[\b[cdeeJvs|$D&ST*[\b[cdeedD!FGGdD!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 /ell /W] / / 3 %(.38;EJ\\
6666)* B! !8C= !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 UORC RC RC RCn !!$H$H$H $H 	$H
 UO$H $H $H 
$Hr!   