Ë
    ÿÍ:jB  ã                   ó"  — d dl Z d dlZd dlZd dlmZmZmZ d dlmZ	 d dl
Zd dlZd dlmZ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!m"Z"m#Z#  e$ejJ                  «      Z& e$ejJ                  «      Z' G d„ de«      Z(y)é    N)ÚDictÚSequenceÚUnion)ÚMLFlowLoggerÚTensorBoardLogger)ÚProblemÚTaskÚ	get_dtype)Úcreate_rng_for_worker)ÚScopeÚSubset)Ú
functional©Údefault_collate)ÚMetric)ÚBinaryAUROCÚMulticlassAUROCÚMultilabelAUROCc                   óö   — e Zd ZdZd„ Zdeeee   ee	ef   f   fd„Z
dej                  fd„Zd„ Zdej                   fd„Zdej                   fd	„Zdej                   fd
„Zdd„Zd„ Zdefd„Zd„ Zd„ Zdefd„Zy)ÚSegmentationTaskz)Methods common to most segmentation tasksc                 ó*   — d| j                   d   |   iS )NÚaudioz
audio-path)Úprepared_data)ÚselfÚfile_ids     ú}/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/pyannote/audio/tasks/segmentation/mixins.pyÚget_filezSegmentationTask.get_file0   s   € Ø˜×+Ñ+¨LÑ9¸'ÑBÐCÐCó    Úreturnc                 óÀ  — t        | j                  j                  «      }| j                  j                  t        j
                  k(  rt        d¬«      S | j                  j                  t        j                  k(  rt        |dd¬«      S | j                  j                  t        j                  k(  rt        |dd¬«      S t        d| j                  j                  › d�«      ‚)z5Returns macro-average of the area under the ROC curveT)Úcompute_on_cpuÚmacro)Úaverager!   zThe zB problem type hasn't been given a default segmentation metric yet.)ÚlenÚspecificationsÚclassesÚproblemr   ÚBINARY_CLASSIFICATIONr   ÚMULTI_LABEL_CLASSIFICATIONr   ÚMONO_LABEL_CLASSIFICATIONr   ÚRuntimeError)r   Únum_classess     r   Údefault_metriczSegmentationTask.default_metric3   sº   € ô
 ˜$×-Ñ-×5Ñ5Ó6ˆØ×Ñ×&Ñ&¬'×*GÑ*GÒGÜ¨dÔ3Ð3Ø× Ñ ×(Ñ(¬G×,NÑ,NÒNÜ" ;¸ÐPTÔUÐUØ× Ñ ×(Ñ(¬G×,MÑ,MÒMÜ" ;¸ÐPTÔUÐUäØ�t×*Ñ*×2Ñ2Ð3Ð3uÐvóð r   Úrngc           	   +   óÒ  K  — | j                   d   d   t        j                  d«      k(  }|j                  «       D ]<  \  }}|| j                   d   |   | j                   d   |   j                  |«      k(  z  }Œ> t	        j
                  |«      d   }| j                   d   |   }t	        j                  |t	        j                  |«      z  «      }| j                  }	t        | dd«      }
	 ||j                  |j                  «       «         }t        |
«      D ]Í  }| j                   d	   |   \  }}t	        j                  | j                   d
   d   || t	        j                  | j                   d
   d   || «      z  «      }||j                  |j                  «       «      z   }| j                   d
   |   \  }}}|j                  |||z   |	z
  «      }| j                  |||	«      –— ŒÏ Œþ­w)aÀ  Iterate over training samples with optional domain filtering

        Parameters
        ----------
        rng : random.Random
            Random number generator
        filters : dict, optional
            When provided (as {key: value} dict), filter training files so that
            only files such as file[key] == value are used for generating chunks.

        Yields
        ------
        chunk : dict
            Training chunks.
        úaudio-metadataÚsubsetÚtrainÚmetadatar   úaudio-annotatedÚnum_chunks_per_fileé   zaudio-regions-idsúannotations-regionsÚduration)r   ÚSubsetsÚindexÚitemsÚnpÚwhereÚcumsumÚsumr8   ÚgetattrÚsearchsortedÚrandomÚrangeÚuniformÚprepare_chunk)r   r.   ÚfiltersÚtrainingÚkeyÚvalueÚfile_idsÚannotated_durationÚcum_prob_annotated_durationr8   r5   r   Ú_Ústart_idÚend_idÚ#cum_prob_annotated_regions_durationÚannotated_region_indexÚregion_durationÚstartÚ
start_times                       r   Útrain__iter__helperz$SegmentationTask.train__iter__helperD   s+  è ø€ ð$ ×%Ñ%Ð&6Ñ7¸ÑAÄWÇ]Á]ØóF
ñ 
ˆð "Ÿ-™-›/ò 	 ‰JˆC�Ø˜×*Ñ*Ð+;Ñ<¸SÑAÀT×EWÑEWØñFàñFç‘5˜“<ñ ñ  ‰Hð	 ô —8‘8˜HÓ% aÑ(ˆð "×/Ñ/Ð0AÑBÀ8ÑLÐÜ&(§i¡iØ¤§¡Ð(:Ó!;Ñ;ó'
Ð#ð —=‘=ˆä% dÐ,AÀ1ÓEÐààÐ:×GÑGÈÏ
É
ËÓUÑVˆGô Ð.Ó/ò H�à#'×#5Ñ#5Ð6IÑ#JÈ7Ñ#SÑ �˜&ô 79·i±iØ×&Ñ&Ð'<Ñ=¸jÑIØ  ðô —f‘fØ×*Ñ*Ð+@ÑAÀ*ÑMØ$ Vðóñó	7Ð3ð Ø9×FÑFÀsÇzÁzÃ|ÓTñUð 'ð -1×,>Ñ,>Ð?TÑ,UØ*ñ-Ñ)��? Eð !Ÿ[™[¨°¸Ñ0GÈ(Ñ0RÓS�
à×(Ñ(¨°*¸hÓGÓGð9Hð ùs   ‚G%G'c              #   óÐ  K  — t        | j                  «      }t        | dd«      }|€| j                  |«      }ntt	        «       }t        j                  |D �cg c]  }| j                  d   |   ‘Œ c}Ž D ]7  }t        ||«      D ��ci c]  \  }}||“Œ
 }}} | j                  |fi |¤Ž||<   Œ9 	 |�|j                  t        |«      «         }t        «      –— Œ-c c}w c c}}w ­w)aV  Iterate over training samples

        Yields
        ------
        dict:
            X: (time, channel)
                Audio chunks.
            y: (frame, )
                Frame-level targets. Note that frame < time.
                `frame` is infered automagically from the
                example model output.
            ...
        ÚbalanceNr3   )r   Úmodelr@   rU   ÚdictÚ	itertoolsÚproductr   ÚzipÚchoiceÚlistÚnext)	r   r.   rW   ÚchunksÚ	subchunksrH   r[   rI   rF   s	            r   Útrain__iter__zSegmentationTask.train__iter__Œ   sý   è ø€ ô  $ D§J¡JÓ/ˆä˜$ 	¨4Ó0ˆØˆ?Ø×-Ñ-¨cÓ2‰Fô ›ˆIÜ$×,Ñ,ØAHÖI¸#�$×$Ñ$ ZÑ0°Ó5ÒIðò N�ô 9<¸GÀWÓ8M×N©*¨#¨u˜3 ™:ÐN�ÑNØ%= T×%=Ñ%=¸cÑ%MÀWÑ%M�	˜'Ò"ðNð ð Ð"Ø" 3§:¡:¬d°9«oÓ#>Ñ?�ô �v“,Òð ùò Jùó
 Oùs   ‚AC&ÁCÁ/C&ÂC ÂAC&c                 ó.  — t        d„ |D «       «      }t        |«      dk(  rt        |D �cg c]  }|d   ‘Œ	 c}«      S t        |«      }t        |D �cg c]0  }t	        j
                  |d   d||d   j                  d   z
  f«      ‘Œ2 c}«      S c c}w c c}w )Nc              3   ó@   K  — | ]  }|d    j                   d   –— Œ y­w)ÚXéÿÿÿÿN)Úshape)Ú.0Úbs     r   ú	<genexpr>z-SegmentationTask.collate_X.<locals>.<genexpr>¸   s   è ø€ Ò6¨1�a˜‘f—l‘l 2Õ&Ñ6ùs   ‚r6   re   r   rf   )Úsetr$   r   ÚmaxÚFÚpadrg   )r   ÚbatchÚlengthsri   Úmax_lens        r   Ú	collate_XzSegmentationTask.collate_X·   s‘   € ÜÑ6°Ô6Ó6ˆô ˆw‹<˜1ÒÜ"°EÖ#:¨q A c£FÒ#:Ó;Ð;ô �g“,ˆÜØEJÖKÀŒQ�U‰U�1�S‘6˜A˜w¨¨3©¯©°bÑ)9Ñ9Ð:Õ;ÒKó
ð 	
ùò	 $;ùò
 Ls   ªBÁ5Bc                 óX   — t        |D �cg c]  }|d   j                  ‘Œ c}«      S c c}w )NÚy)r   Údata©r   ro   ri   s      r   Ú	collate_yzSegmentationTask.collate_yÄ   s#   € Ü°UÖ;°  #¡§£Ò;Ó<Ð<ùÒ;s   Š'c                 óD   — t        |D �cg c]  }|d   ‘Œ	 c}«      S c c}w )NÚmetar   rv   s      r   Úcollate_metazSegmentationTask.collate_metaÇ   s   € Ü°5Ö9¨a  &£	Ò9Ó:Ð:ùÒ9s   Šc                 óz  — | j                  |«      }| j                  |«      }| j                  |«      }| j                  j	                  |dk(  ¬«       | j                  || j
                  j                  j                  |j                  d«      ¬«      }|j                  |j                  j                  d«      |dœS )a£  Collate function used for most segmentation tasks

        This function does the following:
        * stack waveforms into a (batch_size, num_channels, num_samples) tensor batch["X"])
        * apply augmentation when in "train" stage
        * convert targets into a (batch_size, num_frames, num_classes) tensor batch["y"]
        * collate any other keys that might be present in the batch using pytorch default_collate function

        Parameters
        ----------
        batch : list of dict
            List of training samples.

        Returns
        -------
        batch : dict
            Collated batch as {"X": torch.Tensor, "y": torch.Tensor} dict.
        r2   )Úmoder6   )ÚsamplesÚsample_rateÚtargets)re   rt   ry   )rr   rw   rz   Úaugmentationr2   rX   Úhparamsr~   Ú	unsqueezer}   r   Úsqueeze)r   ro   ÚstageÚ
collated_XÚ
collated_yÚcollated_metaÚ	augmenteds          r   Ú
collate_fnzSegmentationTask.collate_fnÊ   s·   € ð* —^‘^ EÓ*ˆ
ð —^‘^ EÓ*ˆ
ð ×)Ñ)¨%Ó0ˆð 	×Ñ×Ñ e¨wÑ&6ÐÔ8Ø×%Ñ%ØØŸ
™
×*Ñ*×6Ñ6Ø×(Ñ(¨Ó+ð &ó 
ˆ	ð ×"Ñ"Ø×"Ñ"×*Ñ*¨1Ó-Ø!ñ
ð 	
r   c                 ó4  — t        j                  | j                  d   d   t        j	                  d«      k(  «      d   }t        j
                  | j                  d   |   «      }t        | j                  t        j                  || j                  z  «      «      S )Nr0   r1   r2   r   r4   )r<   r=   r   r9   r:   r?   rl   Ú
batch_sizeÚmathÚceilr8   )r   Útrain_file_idsr8   s      r   Útrain__len__zSegmentationTask.train__len__õ   s~   € äŸ™Ø×ÑÐ/Ñ0°Ñ:¼g¿m¹mÈGÓ>TÑTó
à
ñˆô —6‘6˜$×,Ñ,Ð->Ñ?ÀÑOÓPˆÜ�4—?‘?¤D§I¡I¨h¸¿¹Ñ.FÓ$GÓHÐHr   r   c                 ó  — t        «       }t        j                  |d   d   t        j	                  d«      k(  «      d   }|D ]x  }|d   |d   d   |k(     }|D ]`  }t        |d   | j                  z  «      }t        |«      D ]5  }|d   || j                  z  z   }	|j                  ||	| j                  f«       Œ7 Œb Œz dt        t        d	„ |D «       «      «      fd
dg}
t        j                  ||
¬«      |d<   |j                  «        y )Nr0   r1   Údevelopmentr   r7   r   r8   rS   c              3   ó&   K  — | ]	  }|d    –— Œ y­w)r   N© )rh   Úvs     r   rj   z6SegmentationTask.prepare_validation.<locals>.<genexpr>  s   è ø€ Ò> q˜a �dÑ>ùs   ‚)rS   Úf)r8   r•   )ÚdtypeÚ
validation)r^   r<   r=   r9   r:   Úroundr8   rC   Úappendr
   rl   ÚarrayÚclear)r   r   Úvalidation_chunksÚvalidation_file_idsr   Úannotated_regionsÚannotated_regionÚ
num_chunksÚcrT   r–   s              r   Úprepare_validationz#SegmentationTask.prepare_validationþ   s=  € Ü ›FÐô !Ÿh™hØÐ*Ñ+¨HÑ5¼¿¹À}Ó9UÑUó
à
ñÐð
 +ò 	SˆGà -Ð.CÑ DØÐ3Ñ4°YÑ?À7ÑJñ!Ðð
 %6ò SÐ ä"Ð#3°JÑ#?À4Ç=Á=Ñ#PÓQ�
ô ˜zÓ*ò S�AØ!1°'Ñ!:¸QÀÇÁÑ=NÑ!N�JØ%×,Ñ,¨g°zÀ4Ç=Á=Ð-QÕRñSñSð	Sð$ Üœ#Ñ>Ð,=Ô>Ó>Ó?ðð Øð
ˆô ')§h¡hÐ/@ÈÔ&Nˆ�lÑ#Ø×ÑÕ!r   c                 ó`   — | j                   d   |   }| j                  |d   |d   |d   ¬«      S )Nr—   r   rS   r8   )r8   )r   rE   )r   ÚidxÚvalidation_chunks      r   Úval__getitem__zSegmentationTask.val__getitem__#  sH   € Ø×-Ñ-¨lÑ;¸CÑ@ÐØ×!Ñ!Ø˜YÑ'Ø˜WÑ%Ø% jÑ1ð "ó 
ð 	
r   c                 ó2   — t        | j                  d   «      S )Nr—   )r$   r   )r   s    r   Ú
val__len__zSegmentationTask.val__len__+  s   € Ü�4×%Ñ% lÑ3Ó4Ð4r   Ú	batch_idxc                 ó(  — |d   |d   }}| j                  |«      }|j                  \  }}}t        | j                  d   | j                  z  |z  «      }t        | j                  d   | j                  z  |z  «      }	|dd…|||	z
  d…f   }
|dd…|||	z
  d…f   }| j
                  j                  t        j                  k(  r;| j                   j                  |
j                  d«      |j                  d«      «       nŸ| j
                  j                  t        j                  k(  rG| j                   j                  t        j                  |
dd«      t        j                  |dd«      «       n1| j
                  j                  t        j                  k(  r
t        «       ‚| j                   j!                  | j                   j                  d	d
d
d
¬«       | j                   j"                  dk(  s4t%        j&                  | j                   j"                  «      dz  dkD  s|dkD  ry|j)                  «       j+                  «       }|j-                  «       j)                  «       j+                  «       }|j)                  «       j+                  «       }t/        | j0                  d«      }t%        j2                  t%        j4                  |«      «      }t%        j2                  ||z  «      }t7        j8                  d|z  |dd	¬«      \  }}t:        j<                  ||dk(  <   t?        |j                  «      dk(  r|dd…dd…t:        j@                  f   }|t;        jB                  |j                  d   «      z  }tE        |«      D �]F  }||z  }||z  }||dz  dz   |f   }||   }|jG                  |«       |jI                  dt?        |«      «       |jK                  d|j                  d   «       |jM                  «       jO                  d	«       |jQ                  «       jO                  d	«       ||dz  dz   |f   }||   }|jS                  d|ddd¬«       |jS                  ||	z
  |ddd¬«       |jG                  |«       |jK                  dd«       |jI                  dt?        |«      «       |jM                  «       jO                  d	«       �ŒI t7        jT                  «        | j                   jV                  D ]•  }tY        |tZ        «      r2|j\                  j_                  d|| j                   j"                  «       ŒEtY        |t`        «      sŒV|j\                  jc                  |jd                  |d| j                   j"                  › d�¬«       Œ— t7        jf                  |«       y)zËCompute validation area under the ROC curve

        Parameters
        ----------
        batch : dict of torch.Tensor
            Current batch.
        batch_idx: int
            Batch index.
        re   rt   r   r6   Né
   rf   é   FT)Úon_stepÚon_epochÚprog_barÚloggeré	   )é   é   )ÚnrowsÚncolsÚfigsizerƒ   Úkg      à?)ÚcolorÚalphaÚlwgš™™™™™¹¿gš™™™™™ñ?r}   Úsamples_epochz.png)Úrun_idÚfigureÚartifact_file)4rX   rg   r˜   Úwarm_upr8   r%   r'   r   r(   Úvalidation_metricÚreshaper)   ÚtorchÚ	transposer*   ÚNotImplementedErrorÚlog_dictÚcurrent_epochrŒ   Úlog2ÚcpuÚnumpyÚfloatÚminr‹   r�   ÚsqrtÚpltÚsubplotsr<   Únanr$   ÚnewaxisÚarangerC   ÚplotÚset_xlimÚset_ylimÚ	get_xaxisÚset_visibleÚ	get_yaxisÚaxvspanÚtight_layoutÚloggersÚ
isinstancer   Ú
experimentÚ
add_figurer   Ú
log_figurer¼   Úclose)r   ro   r©   re   rt   Úy_predrM   Ú
num_framesÚwarm_up_leftÚwarm_up_rightÚpredsÚtargetÚnum_samplesr´   rµ   ÚfigÚaxesÚ
sample_idxÚrow_idxÚcol_idxÚax_refÚsample_yÚax_hypÚsample_y_predr°   s                            r   Úvalidation_stepz SegmentationTask.validation_step.  sº  € ð �S‰z˜5 ™:ˆ1ˆð —‘˜A“ˆØ!Ÿ<™<Ñˆˆ:�qô
 ˜TŸ\™\¨!™_¨t¯}©}Ñ<¸zÑIÓJˆÜ˜dŸl™l¨1™o°·±Ñ=À
ÑJÓKˆØ’q˜,¨°mÑ)CÀbÐHÐHÑIˆØ’1�l Z°-Ñ%?À"ÐDÐDÑEˆð ×Ñ×&Ñ&¬'×*GÑ*GÒGð �J‰J×(Ñ(Ø—‘˜bÓ!Ø—‘˜rÓ"õð
 × Ñ ×(Ñ(¬G×,NÑ,NÒNð �J‰J×(Ñ(Ü—‘  q¨!Ó,Ü—‘ ¨¨1Ó-õð
 × Ñ ×(Ñ(¬G×,MÑ,MÒMä%Ó'Ð'à�
‰
×ÑØ�J‰J×(Ñ(ØØØØð 	ô 	
ð �J‰J×$Ñ$¨Ò)Ü�y‰y˜Ÿ™×1Ñ1Ó2°QÑ6¸Ò:Ø˜1Š}àð �E‰E‹G�M‰M‹OˆØ�G‰G‹I�M‰M‹O×!Ñ!Ó#ˆØ—‘“×#Ñ#Ó%ˆô ˜$Ÿ/™/¨1Ó-ˆÜ—	‘	œ$Ÿ)™) KÓ0Ó1ˆÜ—	‘	˜+¨Ñ-Ó.ˆÜ—L‘LØ�e‘) 5°&À%ô
‰	ˆˆTô
 —F‘Fˆˆ!ˆq‰&‰	Üˆq�w‰w‹<˜1ÒØ’!’QœŸ
™
Ð"Ñ#ˆAØ	ŒR�Y‰Y�q—w‘w˜q‘zÓ"Ñ"ˆô   Ó,ó 	2ˆJà  EÑ)ˆGØ  5Ñ(ˆGð ˜' A™+¨™/¨7Ð2Ñ3ˆFØ˜‘}ˆHØ�K‰K˜Ô!Ø�O‰O˜Aœs 8›}Ô-Ø�O‰O˜B §¡¨qÑ 1Ô2Ø×ÑÓ×*Ñ*¨5Ô1Ø×ÑÓ×*Ñ*¨5Ô1ð ˜' A™+¨™/¨7Ð2Ñ3ˆFØ" :Ñ.ˆMØ�N‰N˜1˜l°#¸SÀQˆNÔGØ�N‰NØ˜]Ñ*¨J¸cÈÐQRð ô ð �K‰K˜Ô&Ø�O‰O˜D #Ô&Ø�O‰O˜Aœs 8›}Ô-Ø×ÑÓ×*Ñ*¨5Ö1ð1	2ô4 	×ÑÔà—j‘j×(Ñ(ò 	ˆFÜ˜&Ô"3Ô4Ø×!Ñ!×,Ñ,¨Y¸¸T¿Z¹Z×=UÑ=UÕVÜ˜F¤LÕ1Ø×!Ñ!×,Ñ,Ø!Ÿ=™=ØØ$1°$·*±*×2JÑ2JÐ1KÈ4Ð"Pð -õ ð		ô 	�	‰	�#�r   N)r2   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   r   r   r   r   Ústrr-   rB   ÚRandomrU   rb   rÂ   ÚTensorrr   rw   rz   r‰   r�   r¢   r¦   r¨   Úintrð   r“   r   r   r   r   -   s®   „ Ù3òDðà	ˆv�x Ñ'¨¨c°6¨kÑ):Ð:Ñ	;óð"FH v§}¡}ó FHòP)ðV
 %§,¡,ó 
ð= %§,¡,ó =ð; U§\¡\ó ;ó)
òVIð#"°ó #"òJ
ò5ðG°ô Gr   r   ))rZ   rŒ   rB   Útypingr   r   r   Úmatplotlib.pyplotÚpyplotrÍ   rÉ   r<   rÂ   Úlightning.pytorch.loggersr   r   Úpyannote.audio.core.taskr   r	   r
   Úpyannote.audio.utils.randomr   Ú#pyannote.database.protocol.protocolr   r   Útorch.nnr   rm   Útorch.utils.data._utils.collater   Útorchmetricsr   Útorchmetrics.classificationr   r   r   r^   Ú__args__r9   ÚScopesr   r“   r   r   ú<module>r     sg   ðó0 Û Û ß (Ñ (å Û Û ß Eß =Ñ =Ý =ß =Ý $Ý ;Ý ß UÑ Uá
ˆv�‰Ó
€Ù	ˆe�n‰nÓ	€ôH�tõ Hr   