a
    d                  
   @   s  d dl mZmZmZmZ d dlmZ d dlmZ d dl	m
Z
mZ d dlmZmZ d dlmZmZ d dlmZ d eeeed
 ed eeeef dddZeeedddZd!eeeedddZd"eeee eeee edddZeeedddZd#eeeedddZdS )$    )AnyCallableOptionalTuple)Tensor)Literal)permutation_invariant_trainingpit_permutate)'scale_invariant_signal_distortion_ratiosignal_distortion_ratio)"scale_invariant_signal_noise_ratiosignal_noise_ratio)_deprecated_root_import_funcspeaker-wisemax)r   zpermutation-wise)r   min)predstargetmetric_funcmode	eval_funckwargsreturnc                 K   s$   t dd tf | ||||d|S )aG  Wrapper for deprecated import.

    >>> from torch import tensor
    >>> preds = tensor([[[-0.0579,  0.3560, -0.9604], [-0.1719,  0.3205,  0.2951]]])
    >>> target = tensor([[[ 1.0958, -0.1648,  0.5228], [-0.4100,  1.1942, -0.5103]]])
    >>> best_metric, best_perm = _permutation_invariant_training(
    ...     preds, target, _scale_invariant_signal_distortion_ratio)
    >>> best_metric
    tensor([-5.1091])
    >>> best_perm
    tensor([[0, 1]])
    >>> pit_permutate(preds, best_perm)
    tensor([[[-0.0579,  0.3560, -0.9604],
             [-0.1719,  0.3205,  0.2951]]])

    r   audio)r   r   r   r   r   )r   r   )r   r   r   r   r   r    r   r/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/torchmetrics/functional/audio/_deprecated.py_permutation_invariant_training   s    

r   )r   permr   c                 C   s   t dd t| |dS )zWrapper for deprecated import.r	   r   r   r   )r   r	   r   r   r   r   _pit_permutate*   s    
r   F)r   r   	zero_meanr   c                 C   s   t dd t| ||dS )zWrapper for deprecated import.

    >>> from torch import tensor
    >>> target = tensor([3.0, -0.5, 2.0, 7.0])
    >>> preds = tensor([2.5, 0.0, 2.0, 8.0])
    >>> _scale_invariant_signal_distortion_ratio(preds, target)
    tensor(18.4030)

    r
   r   r   r   r    )r   r
   r!   r   r   r   (_scale_invariant_signal_distortion_ratio0   s    

r"   N   )r   r   use_cg_iterfilter_lengthr    	load_diagr   c                 C   s   t dd t| |||||dS )a  Wrapper for deprecated import.

    >>> import torch
    >>> g = torch.manual_seed(1)
    >>> preds = torch.randn(8000)
    >>> target = torch.randn(8000)
    >>> _signal_distortion_ratio(preds, target)
    tensor(-12.0589)
    >>> # use with permutation_invariant_training
    >>> preds = torch.randn(4, 2, 8000)  # [batch, spk, time]
    >>> target = torch.randn(4, 2, 8000)
    >>> best_metric, best_perm = _permutation_invariant_training(preds, target, _signal_distortion_ratio)
    >>> best_metric
    tensor([-11.6375, -11.4358, -11.7148, -11.6325])
    >>> best_perm
    tensor([[1, 0],
            [0, 1],
            [1, 0],
            [0, 1]])

    r   r   r   r   r$   r%   r    r&   )r   r   r'   r   r   r   _signal_distortion_ratio>   s    
r(   )r   r   r   c                 C   s   t dd t| |dS )zWrapper for deprecated import.

    >>> from torch import tensor
    >>> target = tensor([3.0, -0.5, 2.0, 7.0])
    >>> preds = tensor([2.5, 0.0, 2.0, 8.0])
    >>> _scale_invariant_signal_noise_ratio(preds, target)
    tensor(15.0918)

    r   r   r   r   )r   r   r)   r   r   r   #_scale_invariant_signal_noise_ratiof   s    

r*   c                 C   s   t dd t| ||dS )zWrapper for deprecated import.

    >>> from torch import tensor
    >>> target = tensor([3.0, -0.5, 2.0, 7.0])
    >>> preds = tensor([2.5, 0.0, 2.0, 8.0])
    >>> _signal_noise_ratio(preds, target)
    tensor(16.1805)

    r   r   r!   )r   r   r!   r   r   r   _signal_noise_ratiot   s    

r+   )r   r   )F)Nr#   FN)F)typingr   r   r   r   Ztorchr   Ztyping_extensionsr   Z!torchmetrics.functional.audio.pitr   r	   Z!torchmetrics.functional.audio.sdrr
   r   Z!torchmetrics.functional.audio.snrr   r   Ztorchmetrics.utilities.printsr   r   r   boolr"   intfloatr(   r*   r+   r   r   r   r   <module>   sD     
    (