
     i!                         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 )PowersetzPowerset 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/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/pyannote/audio/utils/powerset.pyr   zPowerset.__init__0   s[    &(Y(:(:(<O]D,B,B,DQV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   s^     a!2!2Q!67H+E$2B2B,CXN ''K(89  O 8  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9I9IJ
a!2!2Q!67H+E$2B2B,CXN34
K/0a
  O 8
 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    yy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    " "YYx0N"XX0088X2.)) eg 
 ||NDLL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?    & yyLLj$,,..Ar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<<;Q8Q+RT\\%8%8		
 F(($*C*CQ)GH 4<<-7R@		,}<"E !'tag >CCELLQR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   s\    " !&2$""#T%5%5'
"
 **+AB  ,-'
 #"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  $s3x.    " *c * *u|| >.5<< .:ell :$ :5<< :67 7D 7U\\ 7
ell 
u|| 
0+4&+CHo+4	sCx+4Z #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#ryy L#r   