Ë
    ÿÍ:jC  ã                  ó„   — d dl mZ d dlmZmZmZmZ d dlZd dl	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e«      Zy)é    )Úannotations)ÚDictÚOptionalÚSequenceÚUnionN)ÚProtocol)ÚBaseWaveformTransform)ÚMetric)ÚTaské   )Ú)SupervisedRepresentationLearningTaskMixinc                  ój   ‡ — e Zd ZdZ	 	 	 	 	 	 	 	 	 	 d	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 dˆ fd„Zd„ Zˆ xZS )Ú+SupervisedRepresentationLearningWithArcFacea«  Supervised representation learning with ArcFace loss

    Representation learning is the task of ...

    Parameters
    ----------
    protocol : Protocol
        pyannote.database protocol
    duration : float, optional
        Chunks duration in seconds. Defaults to two seconds (2.).
    min_duration : float, optional
        Sample training chunks duration uniformly between `min_duration`
        and `duration`. Defaults to `duration` (i.e. fixed length chunks).
    num_classes_per_batch : int, optional
        Number of classes per batch. Defaults to 32.
    num_chunks_per_class : int, optional
        Number of chunks per class. Defaults to 1.
    margin : float, optional
        Margin. Defaults to 28.6.
    scale : float, optional
        Scale. Defaults to 64.
    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).
    c           
     ó€   •— || _         || _        || _        || _        t        ‰| �  |||| j                  ||	|
|¬«       y )N)ÚdurationÚmin_durationÚ
batch_sizeÚnum_workersÚ
pin_memoryÚaugmentationÚmetric)Únum_chunks_per_classÚnum_classes_per_batchÚmarginÚscaleÚsuperÚ__init__r   )ÚselfÚprotocolr   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/embedding/arcface.pyr   z4SupervisedRepresentationLearningWithArcFace.__init__R   sQ   ø€ ð %9ˆÔ!Ø%:ˆÔ"àˆŒØˆŒ
ä‰ÑØØØ%Ø—‘Ø#Ø!Ø%Øð 	õ 		
ó    c                ó.  — | j                  | j                   j                  «      j                  \  }}t        j                  j                  t        | j                  j                  «      || j                  | j                  ¬«      | j                   _        y )N)r   r   )ÚmodelÚexample_input_arrayÚshapeÚpytorch_metric_learningÚlossesÚArcFaceLossÚlenÚspecificationsÚclassesr   r   Ú	loss_func)r   Ú_Úembedding_sizes      r!   Úsetup_loss_funcz;SupervisedRepresentationLearningWithArcFace.setup_loss_funcr   sm   € à ŸJ™J t§z¡z×'EÑ'EÓF×LÑLÑˆˆ>ä6×=Ñ=×IÑIÜ�×#Ñ#×+Ñ+Ó,ØØ—;‘;Ø—*‘*ð	  Jó  
ˆ�
‰
Õr"   )
Ng       @é    r   gš™™™™™<@g      P@NFNN)r   r   r   zOptional[float]r   Úfloatr   Úintr   r3   r   r2   r   r2   r   zOptional[int]r   Úboolr   zOptional[BaseWaveformTransform]r   z2Union[Metric, Sequence[Metric], Dict[str, Metric]])Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   r0   Ú__classcell__)r    s   @r!   r   r   &   s¡   ø„ ñ"ðV )-ØØ%'Ø$%ØØØ%)Ø Ø8<ØEIð
àð
ð &ð
ð ð	
ð
  #ð
ð "ð
ð ð
ð ð
ð #ð
ð ð
ð 6ð
ð Cõ
ö@	
r"   r   )Ú
__future__r   Útypingr   r   r   r   Úpytorch_metric_learning.lossesr'   Úpyannote.databaser   Ú/torch_audiomentations.core.transforms_interfacer	   Útorchmetricsr
   Úpyannote.audio.core.taskr   Úmixinsr   r   © r"   r!   ú<module>rC      s4   ðõ0 #ç 2Ó 2ã %Ý &Ý QÝ å )å =ôU
Ø-ØõU
r"   