
     i$                     b   d dl Z d dlmZmZ d dlmZmZmZmZ d dl	Z
d dlZd dlZd dlmc mZ d dlmZ d dlmZ eddeed   z  defd	       Zd
 Zd Zej4                  	 	 ddej6                  dej6                  deed   z  dedeej6                  eee      f   f
d       Zej4                  	 	 ddej<                  dej<                  deed   z  dedeej<                  eee      f   f
d       Zdefdede deej6                  ej6                  gej6                  f   de
jB                  fdZ"y)    N)partialsingledispatch)CallableListLiteralTuple)SlidingWindowFeature)linear_sum_assignment	cost_func)msemaereturn_costc                     t               )a  Find cost-minimizing permutation

    Parameters
    ----------
    y1 : np.ndarray or torch.Tensor
        (batch_size, num_samples, num_classes_1)
    y2 : np.ndarray or torch.Tensor
        (num_samples, num_classes_2) or (batch_size, num_samples, num_classes_2)
    cost_func : callable or {"mse", "mae"}, optional
        Can be either "mse" (mean squared error) or "mae" (mean absolute error) or a callable.
        When callable, takes two (num_samples, num_classes) sequences 
        and returns (num_classes, ) pairwise cost.
        Defaults to computing mean squared error ("mse").
    return_cost : bool, optional
        Whether to return cost matrix. Defaults to False.

    Returns
    -------
    permutated_y2 : np.ndarray or torch.Tensor
        (batch_size, num_samples, num_classes_1)
    permutations : list of tuple
        List of permutations so that permutation[i] == j indicates that jth speaker of y2
        should be mapped to ith speaker of y1.  permutation[i] == None when none of y2 speakers
        is mapped to ith speaker of y1.
    cost : np.ndarray or torch.Tensor, optional
        (batch_size, num_classes_1, num_classes_2)
    )	TypeError)y1y2r   r   s       u/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/pyannote/audio/utils/permutation.py	permutater   %   s    : +    c                 \    t        j                  t        j                  | |d      d      S )zCompute class-wise mean-squared error

    Parameters
    ----------
    Y, y : (num_frames, num_classes) torch.tensor

    Returns
    -------
    mse : (num_classes, ) torch.tensor
        Mean-squared error
    none)	reductionr   axis)torchmeanFmse_lossYykwargss      r   mse_cost_funcr#   E   s"     ::ajjA8qAAr   c                 \    t        j                  t        j                  | |z
        d      S )zCompute class-wise mean absolute difference error

    Parameters
    ----------
    Y, y: (num_frames, num_classes) torch.tensor

    Returns
    -------
    mae : (num_classes, ) torch.tensor
        Mean absolute difference error
    r   r   )r   r   absr   s      r   mae_cost_funcr&   T   s"     ::eiiA&Q//r   r   r   returnc                    | j                   \  }}}t        |j                         dk(  r|j                  |dd      }t        |j                         dk7  rd}t        |      |j                   \  }}	}
||k7  s||	k7  r:dt	        | j                          dt	        |j                          d}t        |      |d}g }g }|rg }t        j                  | j                   |j                  |j                  	      }t        t        | |            D ]  \  }\  }}t        j                         5  |dk(  r>|j                  d      |j                  d
      z
  }t        j                  ||z  d      }n|dk(  rN|j                  d      |j                  d
      z
  }t        j                  t        j                  |      d      }nOt        j                  t!        |      D cg c]'  } |||d d ||d
z   f   j                  d|
            ) c}      }d d d        |
|kD  r6t#        j$                  ddd|
|z
  fdt        j&                  |      d
z         }n}d g|z  }t        t)        |j+                                D ]!  \  }}||k  s|||<   |d d |f   ||d d |f<   # |j-                  t	        |             |sj-                  |        |r||t        j                        fS ||fS c c}w # 1 sw Y   xY w)N      zAIncorrect shape: should be (batch_size, num_frames, num_classes).zShape mismatch: z vs. .r   )devicedtype   r   )dimr   constant)shapelenexpand
ValueErrortupler   zerosr-   r.   	enumeratezipno_grad	unsqueezer   r%   stackranger   padmaxr
   cpuappend)r   r   r   r   
batch_sizenum_samplesnum_classes_1msgbatch_size_num_samples_num_classes_2permutationspermutated_y2costsby1_y2_diffcostipadded_costpermutationk1k2s                           r   permutate_torchrV   c   s    .0XX*J]
288}YYz2r*
288}Qo/1xx,K}[ K<$? rxx 1uRXX6GqIo	LMKK"((KM"3r2;/:C ]]_E!}}Q'#--*::zz$+15e#}}Q'#--*::zz%))D/q9{{ "'}!5!5A "#s1a!a%i<'8'?'?M'RS!5  =(%%Aq--78		$!#	K Kf},01BCDFBM!"$B*-ae*aBh' E 	E+./LLI 0L lEKK,>>>,&&; _s   *B3K3,K.	K3.K33K<	c                     t        t        j                  |       t        j                  |      ||      }|r'|\  }}}|j                         ||j                         fS |\  }}|j                         |fS )Nr   r   )r   r   
from_numpynumpy)r   r   r   r   outputrJ   rI   rK   s           r   permutate_numpyr\      s{     	F -3*|U""$lEKKMAA"(M< ,..r   g      ?segmentationsonsetc           
         t        ||      }| j                  }| j                  j                  \  }}}t	        j
                  |j                  |j                  z  dz
        }d|fz  }t        j                         }	t        |       D ]a  \  }
\  }}t        t        d|
|d   z
        t        ||
|d   z   dz               D ]%  }||
k(  r
t        |
|z
  |z  |j                  z  |j                  z        }|dk  r| }||d }| |d||z
  f   }n|d||z
   }| ||df   }t        |t         j"                     ||d      \  }\  }\  }t        |      D ]  \  }}t!        j$                  |dd|f   |kD        }t!        j$                  |dd|f   |kD        }|r|	j'                  |
|f       |r|	j'                  ||f       |sq|st|	j)                  |
|f||f|||f           ( d |	S )	aa  Build permutation graph

    Parameters
    ----------
    segmentations : (num_chunks, num_frames, local_num_speakers)-shaped SlidingWindowFeature
        Raw output of segmentation model.
    onset : float, optionan
        Threshold above which a speaker is considered active. Defaults to 0.5
    cost_func : callable
        Cost function used to find the optimal bijective mapping between speaker activations
        of two overlapping chunks. Expects two (num_frames, num_classes) torch.tensor as input
        and returns cost as a (num_classes, ) torch.tensor. Defaults to mae_cost_func.

    Returns
    -------
    permutation_graph : nx.Graph
        Nodes are (chunk_idx, speaker_idx) tuples.
        An edge between two nodes indicate that those are likely to be the same speaker
        (the lower the value of "cost" attribute, the more likely).
    )r^   r/   r)   r   NTrX   )rP   )r   sliding_windowdatar2   mathfloordurationstepnxGraphr8   r=   r?   minroundr   npnewaxisanyadd_nodeadd_edge)r]   r^   r   chunks
num_chunks
num_frames_max_lookahead	lookaheadpermutation_graphCchunksegmentationcshiftthis_segmentationsthat_segmentationsrS   rP   thisthatthis_is_activethat_is_actives                          r   build_permutation_graphr      s&   4 	/I))F - 2 2 8 8J
AJJv<q@AM]$$I
$-m$<  E<s1a)A,./ZYq\AQTUAU1VWAAv 1q5J.<vNOEqy%1%&%9"%216J
U8J6J3J%K"%12FJ4F%G"%21ef9%=" *3"2::."# 	*&A~w (4
d!#(:1d7(Ce(K!L!#(:1d7(Ce(K!L!%..4y9!%..4y9!n%..D	At94d
3C /  51 X %=R r   )r   F)#rb   	functoolsr   r   typingr   r   r   r   networkxrf   rZ   rj   r   torch.nn.functionalnn
functionalr   pyannote.corer	   scipy.optimizer
   boolr   r#   r&   registerTensorintrV   ndarrayr\   floatrg   r    r   r   <module>r      s  2  - 1 1      . 0 GL,A!A X\  >B0  38	I'I'I' ',//I' 	I'
 5<<eCj))*I' I'X  38	/


/


/ ',/// 	/
 2::tE#J''(/ /0 FSL'LL u||4ellBCL XX	Lr   