
      ip6                     f   d dl Z d dlZd dlm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Zd dlZd dlmZ d dlmZ d dlmZm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mZ d d
lm Z  ddl!m"Z" ddlm#Z# ejH                  jK                  ejH                  jL                          G d d      Z'y)    N)Path)IterableOptionalCallable	GeneratorMappingUnionDict)ExperimentalWarning)
BasePruner)BaseSampler
TPESampler)Trial
FixedTrial)
RDBStorageJournalStorageJournalFileStorage)tqdm)	bayes_mvs   )Pipeline)PipelineInputc                   X   e Zd ZdZ	 	 	 	 	 	 ddedee   dee   deeee	f      deeee
f      dee   d	efd
Zedefd       Zedefd       Zedefd       Z	 ddee   deeef   deegef   fdZ	 	 	 ddee   dededeeef   def
dZ	 	 ddee   dedeeef   deeddf   fdZy)	Optimizerac  Pipeline optimizer

    Parameters
    ----------
    pipeline : `Pipeline`
        Pipeline.
    db : `Path`, optional
        Path to trial database on disk. Use ".sqlite" extension for SQLite
        backend, and ".journal" for Journal backend (prefered for parallel
        optimization).
    study_name : `str`, optional
        Name of study. In case it already exists, study will continue from
        there. # TODO -- generate this automatically
    sampler : `str` or sampler instance, optional
        Algorithm for value suggestion. Must be one of "RandomSampler" or
        "TPESampler", or a sampler instance. Defaults to "TPESampler".
    pruner : `str` or pruner instance, optional
        Algorithm for early pruning of trials. Must be one of "MedianPruner" or
        "SuccessiveHalvingPruner", or a pruner instance.
        Defaults to no pruning.
    seed : `int`, optional
        Seed value for the random number generator of the sampler.
        Defaults to no seed.
    average_case : `bool`, optional
        Optimize for average case (default).
        Set to False to optimize for worst case.
    Npipelinedb
study_namesamplerprunerseedaverage_casec           	      >   || _         || _        |d | _        nt        | j                        j                  }|dk(  r3t        j                  d       t        d| j                         | _        nL|dk(  rt        d| j                         | _        n)|dk(  r$t        t        | j                               | _        || _
        t        |t              r|| _        nKt        |t              r(	  t        t         j"                  |      |      | _        n|t)        |      | _        t        |t*              r|| _        n=t        |t              r&	  t        t         j.                  |             | _        nd | _        t!        j0                  | j                  d	| j                  | j                  | j,                  | j                   j3                         
      | _        || _        y # t$        $ r}	d}
t'        |
      d }	~	ww xY w# t$        $ r}	d}
t'        |
      d }	~	ww xY w)Nz.dbzHStorage with '.db' extension has been deprecated. Use '.sqlite' instead.z
sqlite:///z.sqlitez.journal)r    z8`sampler` must be one of "RandomSampler" or "TPESampler"zC`pruner` must be one of "MedianPruner" or "SuccessiveHalvingPruner"T)r   load_if_existsstorager   r   	direction)r   r   storage_r   suffixwarningswarnr   r   r   r   
isinstancer   r   strgetattroptunasamplersAttributeError
ValueErrorr   r   r   prunerscreate_studyget_directionstudy_r!   )selfr   r   r   r   r   r    r!   	extensionemsgs              p/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/pyannote/pipeline/optimizer.py__init__zOptimizer.__init__S   s    !: DMTWW,,IE!^ !+Zy+A Bi' *Zy+A Bj( ./ATWWI/O P$g{+"DL%&@wv@dK _%40DLfj) DK$&=gfnnf=?
 DK ))MMLL;;mm113
 )= " &P o%& " &[ o%&s0   ,&G# $H #	G>,G99G>	H
HHreturnc                     	 | j                   j                  }|S # t        $ r8 | j                  j	                         dk(  rdnd}|t
        j                  z  }Y |S w xY w)zReturn best loss so farminimizer   )r4   
best_value	Exceptionr   r3   npinf)r5   r?   r%   s      r9   	best_losszOptimizer.best_loss   s_    	,//J   	,"&--"="="?:"MQSUI"RVV+J	,s    =AAc                 x    t        | j                  j                        }| j                  j	                  |      S )zReturn best parameters so fartrial)r   r4   best_paramsr   
parameters)r5   rF   s     r9   rG   zOptimizer.best_params   s0     4;;223}}''e'44    c                 L    | j                   j                  | j                        S )z8Return pipeline instantiated with best parameters so far)r   instantiaterG   )r5   s    r9   best_pipelinezOptimizer.best_pipeline   s     }}(()9)9::rI   inputsshow_progressc                 v     t              t              }dk(  rdddddt        dt        f fd}|S )	a  
        Create objective function used by optuna

        Parameters
        ----------
        inputs : `iterable`
            List of inputs to process.
        show_progress : bool or dict
            Show within-trial progress bar using tqdm progress bar.
            Can also be a **kwarg dict passed to tqdm.

        Returns
        -------
        objective : `callable`
            Callable that takes trial as input and returns correspond loss.
        TzCurrent trialFr   )descleavepositionrF   r;   c                    	 j                   j                         }g }g }j                   j                  j                   j	                  |             }dk7  r't        ddt              i}|j                  d       t              D ]f  \  }}	t        j                         }
t        |	t              r|	j                  di       }ni } ||	fi |}t        j                         }|j                  ||
z
         t        j                         }|$|j                  |	|      }j                  |       nddlm}  ||	d   | ||	      	      }t        j                         }|j                  ||z
         dk7  rj                  d
       j"                  | j%                  |t'        j(                        n
t+        |      |       | j-                         sUt/        j0                          dk7  rj3                          | j5                  dt7        |             | j5                  dt7        |             |Ct        t'        j8                              d
k(  r
|d   x}x}}n0t;        |d      \  \  }\  }}}}n|j=                  d      \  }\  }}j>                  r||S t+        |      S j                   jA                         dk(  r|S |S # t        $ r}d}g }Y d}~d}~ww xY w)zCompute objective value

            Parameter
            ---------
            trial : `Trial`
                Current trial

            Returns
            -------
            loss : `float`
                Loss
            NrE   Ftotalr   pipeline_kwargs)get_annotated
annotation)uemr   processing_timeevaluation_timeg?)alphar=    )!r   
get_metricNotImplementedErrorrK   rH   r   lenupdate	enumeratetimer*   r   getappendlosspyannote.databaserV   r   reportrA   meanabsshould_pruner-   TrialPrunedcloseset_user_attrsumuniquer   confidence_intervalr!   r3   )rF   metricr7   lossesrY   rZ   r   progress_bariinputbefore_processingrU   outputafter_processingbefore_evaluationre   rV   _after_evaluationrh   lower_boundupper_boundrM   r5   rN   s                         r9   	objectivez*Optimizer.get_objective.<locals>.objective   s   113
 !O O }}001I1IPU1I1VWH%#G#f+GG##A& &f-5 %)IIK! eW-&+ii0A2&FO&(O!%;?;#'99; &&'7:K'KL %)IIK! >#==7DMM$'
 @u\2Fe@TUA#'99; &&'7:K'KL E) ''*;;&RWWV_CKQRS%%' ,,..W .Z %""$ 133GH 133GH~ryy()Q.7=ay@D@;?Hc@<6T5K1 4:3M3MTW3M3X00{K  >K v;& ==..0J>  !g ' s   K 	K3$K..K3)listr_   r   float)r5   rM   rN   n_inputsr~   s   ```  r9   get_objectivezOptimizer.get_objective   sK    . fv;D %4uRSTMh	U h	u h	T rI   n_iterations
warm_startc                    d| j                   _        | j                  ||      }|rn| j                   j                  |      }t	        j
                         5  t	        j                  dt               | j                  j                  |       ddd       | j                  j                  ||dd       d| j                   _        | j                  | j                  d	S # 1 sw Y   RxY w)
a  Tune pipeline

        Parameters
        ----------
        inputs : iterable
            List of inputs processed by the pipeline at each iteration.
        n_iterations : int, optional
            Number of iterations. Defaults to 10.
        warm_start : dict, optional
            Nested dictionary of initial parameters used to bootstrap tuning.

        Returns
        -------
        result : dict
            ['loss']
            ['params'] nested dictionary of optimal parameters
        TrN   ignorecategoryNr   n_trialstimeoutn_jobsFre   params)r   trainingr   _flattenr(   catch_warningsfilterwarningsr   r4   enqueue_trialoptimizerC   rG   )r5   rM   r   r   rN   r~   flattened_paramss          r9   tunezOptimizer.tune2  s    4 "&&&v]&K	#}}55jA((*'';NO))*:; + 	YtTUV "'$2B2BCC +*s   7CC'c              #     K   | j                  ||      }	 | j                  }|rn| j
                  j                  |      }t        j                         5  t        j                  dt               | j                  j                  |       ddd       	 d| j
                  _        | j                  j                  |ddd       	 | j                  }| j                  }d| j
                  _        ||d	 b# t        $ r}t        j                  }Y d}~d}~ww xY w# 1 sw Y   xY w# t        $ r
}Y d}~d}~ww xY ww)
a  

        Parameters
        ----------
        inputs : iterable
            List of inputs processed by the pipeline at each iteration.
        warm_start : dict, optional
            Nested dictionary of initial parameters used to bootstrap tuning.

        Yields
        ------
        result : dict
            ['loss']
            ['params'] nested dictionary of optimal parameters
        r   Nr   r   Tr   r   Fr   )r   rC   r0   rA   rB   r   r   r(   r   r   r   r4   r   r   r   rG   )	r5   rM   r   rN   r~   rC   r7   r   rG   s	            r9   	tune_iterzOptimizer.tune_iter^  s    , &&v]&K		I #}}55jA((*'';NO))*:; + %)DMM" KK  QQ O NN	"..
 &+DMM"$<<!   	I	
 +*  si   D>C6 1D>7D9D>D( D>6	D?DD>DD>D%!D>(	D;1D>6D;;D>)NNNNNT)F)
   NT)NT)__name__
__module____qualname____doc__r   r   r   r+   r	   r   r   intboolr:   propertyr   rC   dictrG   rL   r   r   r
   r   r   r   r   r   r   r\   rI   r9   r   r   6   s   > "$(5937"!?)?) TN?) SM	?)
 %[ 012?) sJ/0?) sm?) ?)B 5   5T 5 5
 ;x ; ; ,1G'G T4Z(G 
5'5.	!	GX +/*D'*D *D 	*D
 T4Z(*D 
*D^  +/	3='3= 3= T4Z(	3=
 
4t#	$3=rI   r   )(rb   r(   pathlibr   typingr   r   r   r   r   r	   r
   numpyrA   optuna.loggingr-   optuna.prunersoptuna.samplersoptuna.exceptionsr   r   r   r   optuna.trialr   r   optuna.storagesr   r   r   r   scipy.statsr   r   r   r   loggingset_verbosityWARNINGr   r\   rI   r9   <module>r      sq   <    P P P     1 % 3 * J J  J J !  !   V^^33 4[= [=rI   