Ë
    ÿÍ:jŸ$  ã                   ó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/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/pyannote/audio/utils/permutation.pyÚ	permutater   %   s   € ô: ‹+Ðó    c                 ó\   — t        j                  t        j                  | |d¬«      d¬«      S )zÖCompute 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"   € ô �:‰:”a—j‘j  A°Ô8¸qÔAÐAr   c                 ó\   — t        j                  t        j                  | |z
  «      d¬«      S )zíCompute 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"   € ô �:‰:”e—i‘i  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Ö  € ð .0¯X©XÑ*€J�˜]ä
ˆ2�8‰8ƒ}˜ÒØ�Y‰Y�z 2 rÓ*ˆä
ˆ2�8‰8ƒ}˜ÒØQˆÜ˜‹oÐà/1¯x©xÑ,€K�˜}Ø�[Ò  K°<Ò$?Ø ¤ r§x¡x£Ð 1°´u¸R¿X¹X³Ð6GÀqÐIˆÜ˜‹oÐàÐØˆ	à€LØ€MáØˆä—K‘K §¡°·±À"Ç(Á(ÔK€Mä"¤3 r¨2£;Ó/ó $‰ˆ‰:ˆC�ô �]‰]‹_ñ 	Ø˜EÒ!Ø—}‘} QÓ'¨#¯-©-¸Ó*:Ñ:�Ü—z‘z $¨¡+°1Ô5‘Ø˜eÒ#Ø—}‘} QÓ'¨#¯-©-¸Ó*:Ñ:�Ü—z‘z¤%§)¡)¨D£/°qÔ9‘ä—{‘{ô "' }Ó!5öàñ " # sª1¨a°!°a±%¨i¨<Ñ'8×'?Ñ'?ÀÀMÓ'RÕSòó�÷	ð ˜=Ò(ÜŸ%™%ØØ�A�q˜-¨-Ñ7Ð8ØÜ—	‘	˜$“ !Ñ#ó	‰Kð ˆKà�f˜}Ñ,ˆÜÔ0°·±Ó1BÓCÐDò 	5‰FˆB�Ø�MÓ!Ø"$�˜B‘Ø*-ªa°¨e©*�˜a¢ B˜hÒ'ð	5ð 	×ÑœE +Ó.Ô/ãØ�L‰L˜ÖðI$ñL Ø˜l¬E¯K©K¸Ó,>Ð>Ð>à˜,Ð&Ð&ùò;÷	ð 	ú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Ø×"Ñ"Ó$ l°E·K±K³MÐ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�   Ç   s4  € ô4 ˜	¨Ô/€Ià×)Ñ)€FØ -× 2Ñ 2× 8Ñ 8Ñ€J�
˜AÜ—J‘J˜vŸ™°·±Ñ<¸qÑ@ÓA€MØ�]Ð$Ñ$€IäŸ™›
Ðä$-¨mÓ$<ó 'Ñ ˆÑ ˆE�<Ü”s˜1˜a )¨A¡,Ñ.Ó/´°ZÀÀYÈqÁ\ÑAQÐTUÑAUÓ1VÓWó &	ˆAà�AŠvØô ˜1˜q™5 JÑ.°·±Ñ<¸v¿¹ÑNÓOˆEà�qŠyØ˜�Ø%1°%°&Ð%9Ð"Ø%2°1Ð6J¸
ÀUÑ8JÐ6JÐ3JÑ%KÑ"à%1Ð2F°JÀÑ4FÐ%GÐ"Ø%2°1°e±f°9Ñ%=Ð"ô *3Ø"¤2§:¡:Ñ.Ø"Ø#Ø ô	*Ñ&ˆA‰~�™w ô (¨Ó4ò ‘
��dä!#§¡Ð(:º1¸d¸7Ñ(CÀeÑ(KÓ!L�Ü!#§¡Ð(:º1¸d¸7Ñ(CÀeÑ(KÓ!L�á!Ø%×.Ñ.°°4¨yÔ9á!Ø%×.Ñ.°°4¨yÔ9â!¢nØ%×.Ñ.Ø˜D˜	 A t 9°4¸¸d¸
Ñ3Cð /õ òò1&	ð'ð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ð ñ ¨G°LÑ,AÑ!Að ÐX\ò ó ðò>Bò0ð ×Ñð 38Øñ	I'Ø�‰ðI'à�‰ðI'ð ˜' ,Ñ/Ñ/ðI'ð ð	I'ð
 ˆ5�<‰<˜˜e C™jÑ)Ð)Ñ*òI'ó ðI'ðX ×Ñð 38Øñ	/Ø
�
‰
ð/à
�
‰
ð/ð ˜' ,Ñ/Ñ/ð/ð ð	/ð
 ˆ2�:‰:�t˜E #™JÑ'Ð'Ñ(ò/ó ð/ð0 ØFSñLØ'ðLàðLð ˜Ÿ™ u§|¡|Ð4°e·l±lÐBÑCðLð ‡X�Xô	Lr   