
     iq=                         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                /Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/pyannote/audio/tasks/segmentation/multilabel.pyr'   zMultiLabelSegmentation.__init__\   sx     ($89Z[_`h[iZjjl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<<fG $"#xx	48#"//#E{m8D4D3E FC %S/))EG+u- * "((7CD|eW]]5)|D% #, -/HHWBGG,LM.)"DL !%"#xx	48#"//#E{m8D4D3E FC %S/) #L 1C4E E "//#E{m8D4D3E F((,		-(@'A BC #J!(( &)%6T\\9J%J%JE **51%J3 #@ -/HHT\\,QM.) #%(("#S%67rxx#
 !**; <GW8<#GW$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!'']]**LL
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:))..tU;sQ(()?@56yAWL

 (!EII-+e2Du{{2RS

 zz))..TZZ//888

,W5u{{CekkQTXXJJq"((54<"89@@E	jj*51599=KdR((3:&--c2 ZZ**(TZZ//;;;<

 WWD&&~67 ''
 
 BC!T 34W=
=>!$w 12D E"
E3 )*AecAgou$%"

 +tzz))$,,
s %%&67@8@8L8LM8L#x},8L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\%%ffkk%++.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\%%ffkk%++.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"" c4i()" $s)$	"
 " ueE5L112" (4.)" " " c]" " 45" fhv.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+   