Ë
    ÿÍ:jq=  ã                   óÐ   — d dl Z d dlZd dlmZmZmZmZmZmZm	Z	 d dl
Zd dlZd dlmc mZ d dlmZmZmZ d dlmZ d dlmZmZ d dlmZ d dlmZ d dlmZ d d	l m!Z!  G d
„ de«      Z"y)é    N)ÚDictÚListÚOptionalÚSequenceÚTextÚTupleÚUnion)ÚProblemÚ
ResolutionÚSpecifications)ÚSegmentationTask)ÚSegmentÚSlidingWindowFeature)ÚProtocol)ÚSegmentationProtocol)ÚBaseWaveformTransform)ÚMetricc                   ó*  ‡ — e Zd ZdZ	 	 	 	 	 	 	 	 	 	 	 ddedeeedf      deee      de	dee	e
e	e	f   f   deee      d	ee   d
edee   dedee   deeee   eeef   f   fˆ fd„Zdefd„Zdˆ fd„	Zdede	de	fd„Zdefd„Zdefd„Zed„ «       Zˆ xZS )ÚMultiLabelSegmentationa
  Generic multi-label segmentation

    Multi-label segmentation is the process of detecting temporal intervals
    when a specific audio class is active.

    Example use cases include speaker tracking, gender (male/female)
    classification, or audio event detection.

    Parameters
    ----------
    protocol : Protocol
    cache : str, optional
        As (meta-)data preparation might take a very long time for large datasets,
        it can be cached to disk for later (and faster!) re-use.
        When `cache` does not exist, `Task.prepare_data()` generates training
        and validation metadata from `protocol` and save them to disk.
        When `cache` exists, `Task.prepare_data()` is skipped and (meta)-data
        are loaded from disk. Defaults to a temporary path.
    classes : List[str], optional
        List of classes. Defaults to the list of classes available in the training set.
    duration : float, optional
        Chunks duration. Defaults to 2s.
    warm_up : float or (float, float), optional
        Use that many seconds on the left- and rightmost parts of each chunk
        to warm up the model. While the model does process those left- and right-most
        parts, only the remaining central part of each chunk is used for computing the
        loss during training, and for aggregating scores during inference.
        Defaults to 0. (i.e. no warm-up).
    balance: Sequence[Text], optional
        When provided, training samples are sampled uniformly with respect to these keys.
        For instance, setting `balance` to ["database","subset"] will make sure that each
        database & subset combination will be equally represented in the training samples.
    weight: str, optional
        When provided, use this key to as frame-wise weight in loss function.
    batch_size : int, optional
        Number of training samples per batch. Defaults to 32.
    num_workers : int, optional
        Number of workers used for generating training samples.
        Defaults to multiprocessing.cpu_count() // 2.
    pin_memory : bool, optional
        If True, data loaders will copy tensors into CUDA pinned
        memory before returning them. See pytorch documentation
        for more details. Defaults to False.
    augmentation : BaseWaveformTransform, optional
        torch_audiomentations waveform transform, used by dataloader
        during training.
    metric : optional
        Validation metric(s). Can be anything supported by torchmetrics.MetricCollection.
        Defaults to AUROC (area under the ROC curve).
    NÚprotocolÚcacheÚclassesÚdurationÚwarm_upÚbalanceÚweightÚ
batch_sizeÚnum_workersÚ
pin_memoryÚaugmentationÚmetricc                 ó°   •— t        |t        «      st        dt        |«      › d�«      ‚t        ‰| �  |||||	|
|||¬«	       || _        || _        || _        y )NzHMultiLabelSegmentation task expects a SegmentationProtocol but you gave z. )r   r   r   r   r   r    r!   r   )	Ú
isinstancer   Ú
ValueErrorÚtypeÚsuperÚ__init__r   r   r   )Úselfr   r   r   r   r   r   r   r   r   r   r    r!   Ú	__class__s                €ú�/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/pyannote/audio/tasks/segmentation/multilabel.pyr'   zMultiLabelSegmentation.__init__\   sx   ø€ ô ˜(Ô$8Ô9ÜØZÔ[_Ð`hÓ[iÐZjÐjlÐmóð ô 	‰ÑØØØØ!Ø#Ø!Ø%ØØð 	ô 
	
ð ˆŒØˆŒØˆ�ó    Úprepared_datac                 ór  — | j                   €!| j                  st        j                  d«      }| j                  rGt        j                  | j                  j                  «       | j                  j                  «       «      }n| j                  j                  «       }| j                   €Øt        «       }t        «       }|D ]Ž  }|j                  dd «      }|s-t        j                  d|d   › d|d   › d�«      }t        |«      ‚|D ]  }||vsŒ|j                  |«       Œ |j                  |D �cg c]  }|j                  |«      ‘Œ c}«       Œ� t        j                   |t        j"                  ¬«      |d	<   || _         �n=t        «       }|D ]ü  }|j                  dd «      }|s-t        j                  d|d   › d|d   › d�«      }t        |«      ‚t%        |«      t%        | j                   «      z
  }	|	r?t        j                  d|d   › d|d   › d
dj'                  |	«      › d�«      }t)        |«       |j                  t%        |«      t%        | j                   «      z  D �cg c]  }| j                   j                  |«      ‘Œ c}«       Œþ t        j                   | j                   t        j"                  ¬«      |d	<   t        j*                  t-        |«      t-        | j                   «      ft        j.                  ¬«      }
t1        |«      D ]  \  }}d|
||f<   Œ |
|d<   |j3                  «        y c c}w c c}w )NaE  
                Could not infer list of classes. Either provide a list of classes when
                instantiating the task, or make sure that the training protocol provides
                a 'classes' entry. See https://github.com/pyannote/pyannote-database#segmentation
                for more details.
                r   z
                        File "Úuriz" (from ÚdatabaseaW   database) does not
                        provide a 'classes' entry. Please make sure the corresponding
                        training protocol provides a 'classes' entry for all files. See
                        https://github.com/pyannote/pyannote-database#segmentation for more
                        details.
                        ©Údtypeúclasses-listz; database) provides
                        extra classes (z, z,) that are ignored.
                        Túclasses-annotated)r   Úhas_classesÚtextwrapÚdedentÚhas_validationÚ	itertoolsÚchainr   ÚtrainÚdevelopmentÚlistÚgetr$   ÚappendÚindexÚnpÚarrayÚstr_ÚsetÚjoinÚprintÚzerosÚlenÚbool_Ú	enumerateÚclear)r(   r,   ÚmsgÚ
files_iterr   Úannotated_classesÚfileÚfile_classesÚklassÚextra_classesÚannotated_classes_arrayÚfile_ids               r*   Úpost_prepare_dataz(MultiLabelSegmentation.post_prepare_data„   s  € ð �<‰<Ð¨×(8Ò(8Ü—/‘/ðóˆCð ×ÒÜ"Ÿ™Ø—‘×#Ñ#Ó% t§}¡}×'@Ñ'@Ó'Bó‰Jð Ÿ™×,Ñ,Ó.ˆJà�<‰<ÐÜ“fˆGÜ $£Ðà"ò �Ø#Ÿx™x¨	°4Ó8�á#Ü"Ÿ/™/ðØ# E™{˜m¨8°D¸Ñ4DÐ3Eð Fðó�Cô % S›/Ð)à)ò .�EØ GÒ+ØŸ™ uÕ-ð.ð "×(Ñ(Ø7CÖD¨e�W—]‘] 5Õ)ÒDõð%ô, -/¯H©H°WÄBÇGÁGÔ,LˆM˜.Ñ)Ø"ˆDŽLô !%£ÐØ"ò �Ø#Ÿx™x¨	°4Ó8�á#Ü"Ÿ/™/ðØ# E™{˜m¨8°D¸Ñ4DÐ3Eð Fðó�Cô % S›/Ð)ä # LÓ 1´C¸¿¹Ó4EÑ E�Ù Ü"Ÿ/™/ðØ# E™{˜m¨8°D¸Ñ4DÐ3Eð F(Ø(,¯	©	°-Ó(@Ð'Að Bðó�Cô ˜#”Jà!×(Ñ(ô &)¨Ó%6¼¸T¿\¹\Ó9JÑ%Jöà!ð Ÿ™×*Ñ*¨5Õ1òõð3ô@ -/¯H©H°T·\±\ÌÏÉÔ,QˆM˜.Ñ)ô #%§(¡(ÜÐ"Ó#¤S¨¯©Ó%6Ð7¼r¿x¹xô#
Ðô !*Ð*;Ó <ò 	=ÑˆG�WØ8<Ð# G¨WÐ$4Ò5ð	=à-DˆÐ)Ñ*Ø×ÑÕ!ùòm EùòDs   Ä-L/
É"L4
c                 óÞ   •— t         ‰| �  |«       t        | j                  d   t        j
                  t        j                  | j                  | j                  | j                  ¬«      | _        y )Nr2   )r   ÚproblemÚ
resolutionr   Úmin_durationr   )r&   Úsetupr   r,   r
   ÚMULTI_LABEL_CLASSIFICATIONr   ÚFRAMEr   rX   r   Úspecifications)r(   Ústager)   s     €r*   rY   zMultiLabelSegmentation.setupë   sS   ø€ Ü‰‰�eÔä,Ø×&Ñ& ~Ñ6Ü×6Ñ6Ü!×'Ñ'Ø—]‘]Ø×*Ñ*Ø—L‘Lô
ˆÕr+   rS   Ú
start_timec                 óæ  — | j                  |«      }t        |||z   «      }t        «       }| j                  j                  j                  ||«      \  |d<   }| j                  d   | j                  d   d   |k(     }||d   |j                  k  |d   |j                  kD  z     }	| j                  j                  j                  }
d| j                  j                  j                  z  }t        j                  |	d   |j                  «      |j                  z
  |z
  }t        j                  dt        j                  ||
z  «      «      j                  t         «      }t        j"                  |	d   |j                  «      |j                  z
  |z
  }t        j                  ||
z  «      j                  t         «      }| j                  j%                  t        || j                  j&                  j(                  z  «      «      }t        j*                  |t-        | j                  d   «      ft        j.                  ¬	«       }d|d
d
…| j                  d   |   f<   t1        |||	d   «      D ]  \  }}}d|||dz   …|f<   Œ t3        || j                  j                  | j4                  ¬«      |d<   | j                  d   |   }|j6                  j8                  D �ci c]  }|||   “Œ
 c}|d<   ||d   d<   |S c c}w )aï  Prepare chunk for multi-label segmentation

        Parameters
        ----------
        file_id : int
            File index
        start_time : float
            Chunk start time
        duration : float
            Chunk duration.

        Returns
        -------
        sample : dict
            Dictionary containing the chunk data with the following keys:
            - `X`: waveform
            - `y`: target (see Notes below)
            - `meta`:
                - `database`: database index
                - `file`: file index

        Notes
        -----
        y is a trinary matrix with shape (num_frames, num_classes):
            -  0: class is inactive
            -  1: class is active
            - -1: we have no idea

        ÚXzannotations-segmentsrS   ÚstartÚendg      à?r   r2   r0   Nr3   Úglobal_label_idxé   )ÚlabelsÚyzaudio-metadataÚmetarN   )Úget_filer   ÚdictÚmodelÚaudioÚcropr,   rb   ra   Úreceptive_fieldÚstepr   r@   ÚmaximumÚroundÚastypeÚintÚminimumÚ
num_framesÚhparamsÚsample_rateÚonesrG   Úint8Úzipr   r   r1   Únames)r(   rS   r^   r   rN   ÚchunkÚsampleÚ_ÚannotationsÚchunk_annotationsrn   Úhalfra   Ú	start_idxrb   Úend_idxrt   rf   ÚlabelÚmetadataÚkeys                        r*   Úprepare_chunkz$MultiLabelSegmentation.prepare_chunk÷   sÄ  € ð> �}‰}˜WÓ%ˆä˜
 J°Ñ$9Ó:ˆä“ˆØŸ™×)Ñ)×.Ñ.¨t°UÓ;‰ˆˆs‰�Qà×(Ñ(Ð)?Ñ@Ø×ÑÐ5Ñ6°yÑAÀWÑLñ
ˆð
 (Ø˜Ñ! E§I¡IÑ-°+¸eÑ2DÀuÇ{Á{Ñ2RÑSñ
Ðð
 �z‰z×)Ñ)×.Ñ.ˆØ�T—Z‘Z×/Ñ/×8Ñ8Ñ8ˆä—
‘
Ð,¨WÑ5°u·{±{ÓCÀeÇkÁkÑQÐTXÑXˆÜ—J‘J˜q¤"§(¡(¨5°4©<Ó"8Ó9×@Ñ@ÄÓEˆ	ä�j‰jÐ*¨5Ñ1°5·9±9Ó=ÀÇÁÑKÈdÑRˆÜ—(‘(˜3 ™:Ó&×-Ñ-¬cÓ2ˆð —Z‘Z×*Ñ*Ü�(˜TŸZ™Z×/Ñ/×;Ñ;Ñ;Ó<ó
ˆ
ô �W‰WàÜ�D×&Ñ& ~Ñ6Ó7ðô —'‘'ô
ð 
ˆð BCˆŠ!ˆT×ÑÐ 3Ñ4°WÑ=Ð
=Ñ>Ü!$Ø�wÐ 1Ð2DÑ Eó"
ò 	*ÑˆE�3˜ð )*ˆAˆe�c˜A‘gˆo˜uÐ$Ò%ð	*ô
 +Øˆt�z‰z×)Ñ)°$·,±,ô
ˆˆs‰ð ×%Ñ%Ð&6Ñ7¸Ñ@ˆØ8@¿¹×8LÑ8LÖM°˜#˜x¨™}Ñ,ÒMˆˆv‰Ø!(ˆˆv‰�vÑàˆùò Ns   ËK.Ú	batch_idxc                 óh  — |d   }| j                  |«      }|d   }|j                  |j                  k(  sJ ‚|dk7  }||   }||   }t        j                  ||j	                  t
        j                  «      «      }t        j                  |«      ry | j                   j                  d|dddd¬«       d|iS )	Nr`   rf   éÿÿÿÿz
loss/trainFT©Úon_stepÚon_epochÚprog_barÚloggerÚloss)	rj   ÚshapeÚFÚbinary_cross_entropyr%   ÚtorchÚfloatÚisnanÚlog©r(   Úbatchr‡   r`   Úy_predÚy_trueÚmaskr�   s           r*   Útraining_stepz$MultiLabelSegmentation.training_stepK  s·   € Ø�#‰JˆØ—‘˜A“ˆØ�s‘ˆØ�|‰|˜vŸ|™|Ò+Ð+Ð+ð $ r™\ˆØ˜‘ˆØ˜‘ˆÜ×%Ñ% f¨f¯k©k¼%¿+¹+Ó.FÓGˆô �;‰;�tÔØà�
‰
�‰ØØØØØØð 	ô 	
ð ˜ˆ~Ðr+   c                 ó<  — |d   }| j                  |«      }|d   }|j                  |j                  k(  sJ ‚|dk7  }||   }||   }t        j                  ||j	                  t
        j                  «      «      }| j                   j                  d|dddd¬«       d|iS )	Nr`   rf   r‰   úloss/valFTrŠ   r�   )rj   r�   r‘   r’   r%   r“   r”   r–   r—   s           r*   Úvalidation_stepz&MultiLabelSegmentation.validation_steph  s¦   € Ø�#‰JˆØ—‘˜A“ˆØ�s‘ˆØ�|‰|˜vŸ|™|Ò+Ð+Ð+ð $ r™\ˆØ˜‘ˆØ˜‘ˆÜ×%Ñ% f¨f¯k©k¼%¿+¹+Ó.FÓGˆà�
‰
�‰ØØØØØØð 	ô 	
ð ˜ˆ~Ðr+   c                  ó   — y)a‚  Quantity (and direction) to monitor

        Useful for model checkpointing or early stopping.

        Returns
        -------
        monitor : str
            Name of quantity to monitor.
        mode : {'min', 'max}
            Minimize

        See also
        --------
        lightning.pytorch.callbacks.ModelCheckpoint
        lightning.pytorch.callbacks.EarlyStopping
        )rž   Úmin© )r(   s    r*   Úval_monitorz"MultiLabelSegmentation.val_monitorƒ  s   € ð& !r+   )NNg       @g        NNé    NFNN)N)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   r   r	   Ústrr   r”   r   r   r   rr   Úboolr   r   r   r'   rT   rY   r†   rœ   rŸ   Úpropertyr£   Ú__classcell__)r)   s   @r*   r   r   (   sb  ø„ ñ1ðl -1Ø'+ØØ58Ø,0Ø!%ØØ%)Ø Ø8<ØEIñ"àð"ð ˜˜c 4˜iÑ(Ñ)ð"ð ˜$˜s™)Ñ$ð	"ð
 ð"ð �u˜e E¨5 LÑ1Ð1Ñ2ð"ð ˜( 4™.Ñ)ð"ð ˜‘ð"ð ð"ð ˜c‘]ð"ð ð"ð Ð4Ñ5ð"ð �f˜h vÑ.°°S¸&°[Ñ0AÐAÑBõ"ðPe"¨tó e"õN

ðR Sð R°eð RÀuó Rðh¨có ð:°ó ð6 ñ!ó ô!r+   r   )#r8   r5   Útypingr   r   r   r   r   r   r	   Únumpyr@   r“   Útorch.nn.functionalÚnnÚ
functionalr‘   Úpyannote.audio.core.taskr
   r   r   Ú(pyannote.audio.tasks.segmentation.mixinsr   Úpyannote.corer   r   Úpyannote.databaser   Úpyannote.database.protocolr   Ú/torch_audiomentations.core.transforms_interfacer   Útorchmetricsr   r   r¢   r+   r*   ú<module>r¹      sI   ðó0 Û ß E× EÑ Eã Û ß Ð ß HÑ HÝ Eß 7Ý &Ý ;Ý QÝ ôn!Ð-õ n!r+   