Ë
    ÿÍ:j‹!  ã                   ó€   — d dl mZ d dlmZmZ d dlmZmZ d dlZd dl	m
Z
 d dlm
c mZ  G d„ de
j                  «      Zy)é    )Úcached_property)ÚcombinationsÚpermutations)ÚDictÚTupleNc                   ó°  ‡ — e Zd ZdZdedefˆ fd„Zedeee      fd„«       Z	edefd„«       Z
dej                  fd„Zdej                  fd	„Zdd
ej                  dedej                  fd„Zdd
ej                  dedej                  fd„Zdej                  dej                  fd„Zdeedf   deedf   fd„Zedeeedf   eedf   f   fd„«       Zˆ xZS )ÚPowersetzÏPowerset to multilabel conversion, and back.

    Parameters
    ----------
    num_classes : int
        Number of regular classes.
    max_set_size : int
        Maximum number of classes in each set.
    Únum_classesÚmax_set_sizec                 óÆ   •— t         ‰| �  «        || _        || _        | j	                  d| j                  «       d¬«       | j	                  d| j                  «       d¬«       y )NÚmappingF)Ú
persistentÚcardinality)ÚsuperÚ__init__r
   r   Úregister_bufferÚbuild_mappingÚbuild_cardinality)Úselfr
   r   Ú	__class__s      €úr/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/pyannote/audio/utils/powerset.pyr   zPowerset.__init__0   s[   ø€ Ü‰ÑÔØ&ˆÔØ(ˆÔà×Ñ˜Y¨×(:Ñ(:Ó(<ÈÐÔOØ×Ñ˜]¨D×,BÑ,BÓ,DÐQVÐÕWó    Úreturnc                 óÂ   — g }t        d| j                  dz   «      D ]@  }t        t        | j                  «      |«      D ]  }|j	                  t        |«      «       Œ ŒB |S )a3  List of powerset classes

        e.g. with num_classes = 3 and max_set_size = 2:
        {}, {0}, {1}, {2}, {0, 1}, {0, 2}, {1, 2}

        Returns
        -------
        powerset_classes : list of set[int]
            List of powerset classes, each represented as a set of regular class indices.
        r   é   )Úranger   r   r
   ÚappendÚset)r   Úpowerset_classesÚset_sizeÚcurrent_sets       r   r   zPowerset.powerset_classes9   sg   € ð ÐÜ˜a ×!2Ñ!2°QÑ!6Ó7ò 	:ˆHÜ+¬E°$×2BÑ2BÓ,CÀXÓNò :�Ø ×'Ñ'¬¨KÓ(8Õ9ñ:ð	:ð  Ðr   c                 ó,   — t        | j                  «      S )zNumber of powerset classes)Úlenr   ©r   s    r   Únum_powerset_classeszPowerset.num_powerset_classesK   s   € ô �4×(Ñ(Ó)Ð)r   c                 óú   — t        j                  | j                  | j                  «      }d}t	        d| j
                  dz   «      D ]2  }t        t	        | j                  «      |«      D ]  }d|||f<   |dz  }Œ Œ4 |S )at  Compute powerset to regular mapping

        Returns
        -------
        mapping : (num_powerset_classes, num_classes) torch.Tensor
            mapping[i, j] == 1 if jth regular class is a member of ith powerset class
            mapping[i, j] == 0 otherwise

        Example
        -------
        With num_classes == 3 and max_set_size == 2, returns

            [0, 0, 0]  # none
            [1, 0, 0]  # class #1
            [0, 1, 0]  # class #2
            [0, 0, 1]  # class #3
            [1, 1, 0]  # classes #1 and #2
            [1, 0, 1]  # classes #1 and #3
            [0, 1, 1]  # classes #2 and #3

        r   r   )ÚtorchÚzerosr%   r
   r   r   r   )r   r   Ú
powerset_kr    r!   s        r   r   zPowerset.build_mappingP   s†   € ô, —+‘+˜d×7Ñ7¸×9IÑ9IÓJˆØˆ
Ü˜a ×!2Ñ!2°QÑ!6Ó7ò 	 ˆHÜ+¬E°$×2BÑ2BÓ,CÀXÓNò  �Ø34�˜
 KÐ/Ñ0Ø˜a‘‘
ñ ð	 ð
 ˆr   c                 óD   — t        j                  | j                  d¬«      S )z#Compute size of each powerset classr   ©Údim)r'   Úsumr   r$   s    r   r   zPowerset.build_cardinalityo   s   € ä�y‰y˜Ÿ™¨1Ô-Ð-r   ÚpowersetÚsoftc                 ó   — |rt        j                  |«      }nWt         j                  j                  j	                  t        j
                  |d¬«      | j                  «      j                  «       }t        j                  || j                  «      S )a/  Convert predictions from powerset to multi-label

        Parameter
        ---------
        powerset : (batch_size, num_frames, num_powerset_classes) torch.Tensor
            Soft predictions in "powerset" space.
        soft : bool, optional
            Return soft multi-label predictions. Defaults to False (i.e. hard predictions)
            Assumes that `powerset` are "log probabilities".

        Returns
        -------
        multi_label : (batch_size, num_frames, num_classes) torch.Tensor
            Predictions in "multi-label" space.
        éÿÿÿÿr+   )
r'   ÚexpÚnnÚ
functionalÚone_hotÚargmaxr%   ÚfloatÚmatmulr   )r   r.   r/   Úpowerset_probss       r   Úto_multilabelzPowerset.to_multilabels   si   € ñ" Ü"ŸY™Y xÓ0‰Nä"ŸX™X×0Ñ0×8Ñ8Ü—‘˜X¨2Ô.Ø×)Ñ)ó÷ ‰e‹gð ô
 �|‰|˜N¨D¯L©LÓ9Ð9r   c                 ó(   — | j                  ||¬«      S )zAlias for `to_multilabel`)r/   )r:   )r   r.   r/   s      r   ÚforwardzPowerset.forwardŽ   s   € à×!Ñ! (°Ð!Ó6Ð6r   Ú
multilabelc                 ó¾   — t        j                  t        j                  t        j                  || j
                  j                  «      d¬«      | j                  ¬«      S )a›  Convert (hard) predictions from multi-label to powerset

        Parameter
        ---------
        multi_label : (batch_size, num_frames, num_classes) torch.Tensor
            Prediction in "multi-label" space.

        Returns
        -------
        powerset : (batch_size, num_frames, num_powerset_classes) torch.Tensor
            Hard, one-hot prediction in "powerset" space.

        Note
        ----
        This method will not complain if `multilabel` is provided a soft predictions
        (e.g. the output of a sigmoid-ed classifier). However, in that particular
        case, the resulting powerset output will most likely not make much sense.
        r1   r+   )r
   )ÚFr5   r'   r6   r8   r   ÚTr%   )r   r=   s     r   Úto_powersetzPowerset.to_powerset’   s?   € ô& �y‰yÜ�L‰LœŸ™ j°$·,±,·.±.ÓAÀrÔJØ×1Ñ1ô
ð 	
r   Úmultilabel_permutation.c                 óø  — | j                   dd…|f   }t        j                  | j                  | j                   j                  t        j
                  ¬«      }d|z  j                  | j                  df«      }t        j                  | j                   |z  d¬«      }t        j                  ||z  d¬«      }|d   |dd…df   k(  j                  «       j                  d¬«      }t        |j                  «       «      S )a×  Helper function for `permutation_mapping` property

        Takes a (num_classes,)-shaped permutation in multilabel space and returns
        the corresponding (num_powerset_classes,)-shaped permutation in powerset space.
        This does not cache anything and only works on one single permutation at a time.

        Parameters
        ----------
        multilabel_permutation : tuple of int
            Permutation in multilabel space.

        Returns
        -------
        powerset_permutation : tuple of int
            Permutation in powerset space.

        Example
        -------
        >>> powerset = Powerset(3, 2)
        >>> powerset._permutation_powerset((1, 0, 2))
        # (0, 2, 1, 3, 4, 6, 5)

        N)ÚdeviceÚdtypeé   r   r1   r+   r   )r   r'   Úaranger
   rD   ÚintÚtiler%   r-   r6   ÚtupleÚtolist)r   rB   Úpermutated_mappingrG   Úpowers_of_twoÚbeforeÚafterÚpowerset_permutations           r   Ú_permutation_powersetzPowerset._permutation_powersetª   sÜ   € ð6 ,0¯<©<ºÐ;QÐ8QÑ+RÐä—‘Ø×Ñ T§\¡\×%8Ñ%8ÄÇ	Á	ô
ˆð ˜F™×(Ñ(¨$×*CÑ*CÀQÐ)GÓHˆô —‘˜4Ÿ<™<¨-Ñ7¸RÔ@ˆÜ—	‘	Ð,¨}Ñ<À"ÔEˆð !' t¡°²a¸°g±Ñ >×CÑCÓE×LÑLÐQRÐLÓSÐô Ð)×0Ñ0Ó2Ó3Ð3r   c                 ó    — i }t        t        | j                  «      | j                  «      D ]  }| j                  |«      |t	        |«      <   Œ! |S )aÃ  Mapping between multilabel and powerset permutations

        Example
        -------
        With num_classes == 3 and max_set_size == 2, returns

        {
            (0, 1, 2): (0, 1, 2, 3, 4, 5, 6),
            (0, 2, 1): (0, 1, 3, 2, 5, 4, 6),
            (1, 0, 2): (0, 2, 1, 3, 4, 6, 5),
            (1, 2, 0): (0, 2, 3, 1, 6, 4, 5),
            (2, 0, 1): (0, 3, 1, 2, 5, 6, 4),
            (2, 1, 0): (0, 3, 2, 1, 6, 5, 4)
        }
        )r   r   r
   rQ   rJ   )r   Úpermutation_mappingrB   s      r   rS   zPowerset.permutation_mapping×   sc   € ð" !Ðä&2Ü�$×"Ñ"Ó# T×%5Ñ%5ó'
ò 	CÐ"ð
 ×*Ñ*Ð+AÓBð  ÜÐ,Ó-òð	Cð #Ð"r   )F)Ú__name__Ú
__module__Ú__qualname__Ú__doc__rH   r   r   Úlistr   r   r%   r'   ÚTensorr   r   Úboolr:   r<   rA   r   rQ   r   rS   Ú__classcell__)r   s   @r   r	   r	   %   s6  ø„ ñðX Cð X°sõ Xð ð  $ s¨3¡x¡.ò  ó ð ð" ð* cò *ó ð*ð˜uŸ|™|ó ð>. 5§<¡<ó .ñ: e§l¡lð :¸$ð :È5Ï<É<ó :ñ67 §¡ð 7°Dð 7ÀUÇ\Á\ó 7ð
 e§l¡lð 
°u·|±|ó 
ð0+4Ø&+¨C°¨H¡oð+4à	ˆs�Cˆx‰ó+4ðZ ð# T¨%°°S°©/¸5ÀÀcÀ¹?Ð*JÑ%Kò #ó ô#r   r	   )Ú	functoolsr   Ú	itertoolsr   r   Útypingr   r   r'   Útorch.nnr3   Útorch.nn.functionalr4   r?   ÚModuler	   © r   r   ú<module>rc      s.   ðõ8 &ß 0ß ã Ý ß Ð ôL#ˆr�y‰yõ L#r   