
     i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     }/Volumes/fast/ai/experiments/voice-extract-mac/.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33u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AW]]F
 
 "--/JC**+;<SATEWEWFF5<   H * 88H%a( "//0AB8L&(ii(:!;;'
# ==%d,A1E:GG

UVG ./#'#5#56I#J7#S & 79ii&&'<=jI  ff**+@A*M$V	73 9FFszz|TU ' -1,>,>?T,U*-)?E ![[0G(0RS
((*hGG9 0 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       $DJJ/$	40?--c2F I$,,AHI#$$$Z05I 9<GW8MN8M*#u3:8MN%=T%=%=c%MW%M	'"  ""3::d9o#>? 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)XN)shape).0bs     r   	<genexpr>z-SegmentationTask.collate_X.<locals>.<genexpr>   s     61afll2&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#:EqAcFE#:;; g,EJKUQUU1S6Aw3b)99:;U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;U#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95a&	5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 	ew&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:gmmG>TT

 66$,,->?OP4??DIIh.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     >,=qa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 !hh*+H5}9UU


 +G -.C D34Y?7J!
 %6 "#3J#?4==#PQ
 z*A!1'!:Q=N!NJ%,,gz4==-QR + %6 +$ #>,=>>? 
 ')hh/@&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)zCompute 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lwgg?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    Sz5:1 A!<<:q
 T\\!_t}}<zIJdll1o=
JKq,m)CbHHI1lZ-%?"DDE &&'*G*GG JJ((b!r"
   ((G,N,NN JJ((q!,1-
   ((G,M,MM%''

JJ(( 	 	
 JJ$$)yy112Q6:1} EEGMMOGGIMMO!!###% $//1-		$))K01		+-.LLe)5&%
	T
 FF!q&	qww<1!Q

"#A	RYYqwwqz""  ,J E)G 5(G 'A+/723F}HKK!OOAs8}-OOBq 12**51**51 'A+/723F":.MNN1l#SQNGNN]*JcQR   KK&OOD#&OOAs8}-**511 -4 	jj((F&"34!!,,YTZZ=U=UVFL1!!,,!==$1$**2J2J1K4"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	vx'c6k)::	;"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
	enn	Ht Hr   