Ë
    þÍ: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
 d dlmZ d dlmZmZmZ d dlmZmZ dd	d
giZesdgZ G d„ de«      Zy)é    )ÚSequence)ÚAnyÚOptionalÚUnion)ÚTensorÚtensor)Ú'non_intrusive_speech_quality_assessment)ÚMetric)Ú_LIBROSA_AVAILABLEÚ_MATPLOTLIB_AVAILABLEÚ_REQUESTS_AVAILABLE)Ú_AX_TYPEÚ_PLOT_OUT_TYPEÚ#NonIntrusiveSpeechQualityAssessmentÚlibrosaÚrequestsz(NonIntrusiveSpeechQualityAssessment.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d
<   dZeed<   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   uJ  `Non-Intrusive Speech Quality Assessment`_ (NISQA v2.0) [1], [2].

    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

    - ``nisqa`` (:class:`~torch.Tensor`): float tensor reduced across the batch with shape ``(5,)`` corresponding to
      overall MOS, noisiness, discontinuity, coloration and loudness in that order

    .. hint::
        Using this metric requires you to have ``librosa`` and ``requests`` installed. Install as
        ``pip install librosa requests``.

    .. caution::
        The ``forward`` and ``compute`` methods in this class return values reduced across the batch. To obtain
        values for each sample, you may use the functional counterpart
        :func:`~torchmetrics.functional.audio.nisqa.non_intrusive_speech_quality_assessment`.

    Args:
        fs: sampling frequency of input

    Raises:
        ModuleNotFoundError:
            If ``librosa`` or ``requests`` are not installed

    Example:
        >>> import torch
        >>> from torchmetrics.audio import NonIntrusiveSpeechQualityAssessment
        >>> _ = torch.manual_seed(42)
        >>> preds = torch.randn(16000)
        >>> nisqa = NonIntrusiveSpeechQualityAssessment(16000)
        >>> nisqa(preds)
        tensor([1.0433, 1.9545, 2.6087, 1.3460, 1.7117])

    References:
        - [1] G. Mittag and S. MÃ¶ller, "Non-intrusive speech quality assessment for super-wideband speech communication
          networks", in Proc. ICASSP, 2019.
        - [2] G. Mittag, B. Naderi, A. Chehadi and S. MÃ¶ller, "NISQA: A deep CNN-self-attention model for
          multidimensional speech quality prediction with crowdsourced datasets", in Proc. INTERSPEECH, 2021.

    Ú	sum_nisqaÚtotalFÚfull_state_updateÚis_differentiableTÚhigher_is_betterç        Úplot_lower_boundg      @Úplot_upper_boundÚfsÚkwargsÚreturnNc                 ó  •— t        ‰| �  d	i |¤Ž t        rt        st	        d«      ‚t        |t        «      r|dk  rt        d|› �«      ‚|| _        | j                  dt        g d¢«      d¬«       | j                  dt        d«      d¬«       y )
NziNISQA metric requires that librosa and requests are installed. Install as `pip install librosa requests`.r   z9Argument `fs` expected to be a positive integer, but got r   )r   r   r   r   r   Úsum)ÚdefaultÚdist_reduce_fxr   © )ÚsuperÚ__init__r   r   ÚModuleNotFoundErrorÚ
isinstanceÚintÚ
ValueErrorr   Ú	add_stater   )Úselfr   r   Ú	__class__s      €úm/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchmetrics/audio/nisqa.pyr%   z,NonIntrusiveSpeechQualityAssessment.__init__X   sˆ   ø€ Ü‰ÑÑ"˜6Ò"Ý!Õ)<Ü%ð=óð ô ˜"œcÔ" b¨A¢gÜÐXÐY[ÐX\Ð]Ó^Ð^ØˆŒà�‰�{¬FÒ3LÓ,MÐ^cˆÔdØ�‰�w¬¨q«	À%ˆÕHó    Úpredsc                 ó(  — t        || j                  «      j                  | j                  j                  «      }|j                  dd«      }| xj                  |j                  d¬«      z  c_        | xj                  |j                  d   z  c_        y)zUpdate state with predictions.éÿÿÿÿé   r   )ÚdimN)	r	   r   Útor   ÚdeviceÚreshaper    r   Úshape)r+   r/   Únisqa_batchs      r-   Úupdatez*NonIntrusiveSpeechQualityAssessment.updatef   su   € ä=ØØ�G‰Gó
÷ ‰"ˆT�^‰^×"Ñ"Ó
#ð 	ð
 "×)Ñ)¨"¨aÓ0ˆØ�Š˜+Ÿ/™/¨a˜/Ó0Ñ0�Ø�
Š
�k×'Ñ'¨Ñ*Ñ*Ž
r.   c                 ó4   — | j                   | j                  z  S )zCompute metric.)r   r   )r+   s    r-   Úcomputez+NonIntrusiveSpeechQualityAssessment.computeq   s   € à�~‰~ §
¡
Ñ*Ð*r.   ÚvalÚaxc                 ó&   — | j                  ||«      S )aF  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: A 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 NonIntrusiveSpeechQualityAssessment
            >>> metric = NonIntrusiveSpeechQualityAssessment(16000)
            >>> metric.update(torch.randn(16000))
            >>> fig_, ax_ = metric.plot()

        .. plot::
            :scale: 75

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

        )Ú_plot)r+   r<   r=   s      r-   Úplotz(NonIntrusiveSpeechQualityAssessment.plotu   s   € ðL �z‰z˜#˜rÓ"Ð"r.   )NN)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   Ú__annotations__r   Úboolr   r   r   Úfloatr   r(   r   r%   r9   r;   r   r   r   r   r   r@   Ú__classcell__)r,   s   @r-   r   r   #   s¼   ø… ñ*ðX ÓØƒMØ#Ð�tÓ#Ø#Ð�tÓ#Ø!Ð�dÓ!Ø!Ð�eÓ!Ø!Ð�eÓ!ðI˜3ð I¨#ð I°$õ Ið	+˜Fð 	+ tó 	+ð+˜ó +ñ&#˜˜f h¨vÑ&6¸Ð<Ñ=ð &#È(ÐS[ÑJ\ð &#Ðhv÷ &#r.   N)Úcollections.abcr   Útypingr   r   r   Útorchr   r   Ú#torchmetrics.functional.audio.nisqar	   Útorchmetrics.metricr
   Útorchmetrics.utilities.importsr   r   r   Útorchmetrics.utilities.plotr   r   Ú__doctest_requires__Ú__doctest_skip__r   r#   r.   r-   ú<module>rR      sS   ðõ %ß 'Ñ 'ç  å WÝ &÷ñ ÷
 Aà=À	È:Ð?VÐWÐ áØBÐCÐôx#¨&õ x#r.   