Ë
    þÍ:jC  ã                  óÖ   — d dl mZ d dlZd dlZd dlmZ d dlZd dlm	Z	 d dl
mZ d dlmZ er
d dlmZ d dlZ	 	 	 	 	 	 dd„Z	 	 	 	 	 	 	 	 	 	 	 	 dd	„Z	 	 	 	 	 	 	 	 	 	 dd
„Z G d„ de	«      Zy)é    )ÚannotationsN)ÚTYPE_CHECKING)Ú
BasePruner)ÚStudyDirection)Ú
TrialState)ÚKeysViewc                óô   — t        j                  t        | j                  j	                  «       «      t
        ¬«      }|t        j                  k(  rt        j                  |«      S t        j                  |«      S )N©Údtype)
ÚnpÚasarrayÚlistÚintermediate_valuesÚvaluesÚfloatr   ÚMAXIMIZEÚnanmaxÚnanmin)ÚtrialÚ	directionr   s      úo/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/optuna/pruners/_percentile.pyÚ(_get_best_intermediate_result_over_stepsr      sT   € ô �Z‰Zœ˜U×6Ñ6×=Ñ=Ó?Ó@ÌÔN€FØ”N×+Ñ+Ò+Ü�y‰y˜Ó Ð Ü�9‰9�VÓÐó    c                óp  — t        | «      dk(  rt        d«      ‚| D �cg c]   }||j                  v sŒ|j                  |   ‘Œ" }}t        |«      |k  rt        j                  S |t
        j                  k(  rd|z
  }t        t        j                  t        j                  |t        ¬«      |«      «      S c c}w )Nr   zNo trials have been completed.éd   r
   )ÚlenÚ
ValueErrorr   ÚmathÚnanr   r   r   r   ÚnanpercentileÚarray)Úcompleted_trialsr   ÚstepÚ
percentileÚn_min_trialsÚtr   s          r   Ú/_get_percentile_intermediate_result_over_trialsr'      s¶   € ô ÐÓ Ò!ÜÐ9Ó:Ð:ð .>öØ()ÀÈ×I^ÑI^ÒA^ˆ×Ñ˜dÓ#ðÐð ô ÐÓ ,Ò.Ü�x‰xˆà”N×+Ñ+Ò+Ø˜:Ñ%ˆ
äÜ
×ÑÜ�H‰HÐ(´Ô6Øó	
óð ùòs
   žB3²B3c                ól   ‡ — ‰ |z
  |z  |z  |z   }|dk\  sJ ‚t        j                  ˆ fd„|d«      }||k  S )Nr   c                ó    •— || kD  r|‰k7  r|S | S )N© )Úsecond_last_stepÚsr#   s     €r   ú<lambda>z,_is_first_in_interval_step.<locals>.<lambda>C   s   ø€ ¨Ð-=Ò)=À!ÀtÂ) A€ ÐQa€ r   éÿÿÿÿ)Ú	functoolsÚreduce)r#   Úintermediate_stepsÚn_warmup_stepsÚinterval_stepsÚnearest_lower_pruning_stepr+   s   `     r   Ú_is_first_in_interval_stepr5   9   sb   ø€ ð 	ˆ~ÑØ	ñ"à(ñ")à+9ñ":Ðð &¨Ò*Ð*Ð*ô !×'Ñ'ÛaØØ
óÐð Ð8Ñ8Ð8r   c                  óD   — e Zd ZdZ	 	 	 dddœ	 	 	 	 	 	 	 	 	 	 	 dd„Zd	d„Zy)
ÚPercentilePruneraI
  Pruner to keep the specified percentile of the trials.

    Prune if the best intermediate value is in the bottom percentile among trials at the same step.

    Example:

        .. testcode::

            import numpy as np
            from sklearn.datasets import load_iris
            from sklearn.linear_model import SGDClassifier
            from sklearn.model_selection import train_test_split

            import optuna

            X, y = load_iris(return_X_y=True)
            X_train, X_valid, y_train, y_valid = train_test_split(X, y)
            classes = np.unique(y)


            def objective(trial):
                alpha = trial.suggest_float("alpha", 0.0, 1.0)
                clf = SGDClassifier(alpha=alpha)
                n_train_iter = 100

                for step in range(n_train_iter):
                    clf.partial_fit(X_train, y_train, classes=classes)

                    intermediate_value = clf.score(X_valid, y_valid)
                    trial.report(intermediate_value, step)

                    if trial.should_prune():
                        raise optuna.TrialPruned()

                return clf.score(X_valid, y_valid)


            study = optuna.create_study(
                direction="maximize",
                pruner=optuna.pruners.PercentilePruner(
                    25.0, n_startup_trials=5, n_warmup_steps=30, interval_steps=10
                ),
            )
            study.optimize(objective, n_trials=20)

    Args:
        percentile:
            Percentile which must be between 0 and 100 inclusive
            (e.g., When given 25.0, top of 25th percentile trials are kept).
        n_startup_trials:
            Pruning is disabled until the given number of trials finish in the same study.
        n_warmup_steps:
            Pruning is disabled until the trial exceeds the given number of step. Note that
            this feature assumes that ``step`` starts at zero.
        interval_steps:
            Interval in number of steps between the pruning checks, offset by the warmup steps.
            If no value has been reported at the time of a pruning check, that particular check
            will be postponed until a value is reported. Value must be at least 1.
        n_min_trials:
            Minimum number of reported trial results at a step to judge whether to prune.
            If the number of reported intermediate values from all trials at the current step
            is less than ``n_min_trials``, the trial will not be pruned. This can be used to ensure
            that a minimum number of trials are run to completion without being pruned.
    é   )r%   c               ó"  — d|cxk  rdk  sn t        d|›d�«      ‚|dk  rt        d|›d�«      ‚|dk  rt        d|›d�«      ‚|dk  rt        d	|›d�«      ‚|dk  rt        d
|›d�«      ‚|| _        || _        || _        || _        || _        y )Ng        r   zCPercentile must be between 0 and 100 inclusive, but got percentile=ú.r   zFNumber of startup trials cannot be negative, but got n_startup_trials=zBNumber of warmup steps cannot be negative, but got n_warmup_steps=r8   zBPruning interval steps must be at least 1, but got interval_steps=zFNumber of trials for pruning must be at least 1, but got n_min_trials=)r   Ú_percentileÚ_n_startup_trialsÚ_n_warmup_stepsÚ_interval_stepsÚ_n_min_trials)Úselfr$   Ún_startup_trialsr2   r3   r%   s         r   Ú__init__zPercentilePruner.__init__�   sé   € ð �jÔ' CÔ'ÜØVÈ:È-ÐWXÐYóð ð ˜aÒÜØYÐHXÐGZÐZ[Ð\óð ð ˜AÒÜØUÀnÐEVÐVWÐXóð ð ˜AÒÜØUÀnÐEVÐVWÐXóð ð ˜!ÒÜØYÈLÈ?ÐZ[Ð\óð ð &ˆÔØ!1ˆÔØ-ˆÔØ-ˆÔØ)ˆÕr   c                ó4  — |j                  dt        j                  f¬«      }t        |«      }|dk(  ry|| j                  k  ry|j
                  }|€y| j                  }||k  ryt        ||j                  j                  «       || j                  «      sy|j                  }t        ||«      }t        j                  |«      ryt        |||| j                   | j"                  «      }	t        j                  |	«      ry|t$        j&                  k(  r||	k  S ||	kD  S )NF)ÚdeepcopyÚstatesr   T)Ú
get_trialsr   ÚCOMPLETEr   r<   Ú	last_stepr=   r5   r   Úkeysr>   r   r   r   Úisnanr'   r;   r?   r   r   )
r@   Ústudyr   r"   Ún_trialsr#   r2   r   Úbest_intermediate_resultÚps
             r   ÚprunezPercentilePruner.prune±   s  € Ø ×+Ñ+°UÄJ×DWÑDWÐCYÐ+ÓZÐÜÐ'Ó(ˆà�qŠ=Øà�d×,Ñ,Ò,Øà�‰ˆØˆ<Øà×-Ñ-ˆØ�.Ò Øä)Ø�%×+Ñ+×0Ñ0Ó2°NÀD×DXÑDXô
ð à—O‘Oˆ	Ü#KÈEÐS\Ó#]Ð Ü�:‰:Ð.Ô/Øä;Ø˜i¨¨t×/?Ñ/?À×ASÑASó
ˆô �:‰:�aŒ=Øàœ×/Ñ/Ò/Ø+¨aÑ/Ð/Ø'¨!Ñ+Ð+r   N)é   r   r8   )r$   r   rA   Úintr2   rQ   r3   rQ   r%   rQ   ÚreturnÚNone)rK   z'optuna.study.Study'r   ú'optuna.trial.FrozenTrial'rR   Úbool)Ú__name__Ú
__module__Ú__qualname__Ú__doc__rB   rO   r*   r   r   r7   r7   K   sb   „ ñ?ðH !"ØØð"*ð ñ"*àð"*ð ð"*ð ð	"*ð
 ð"*ð ð"*ð 
ó"*ôH$,r   r7   )r   rT   r   r   rR   r   )r"   z list['optuna.trial.FrozenTrial']r   r   r#   rQ   r$   r   r%   rQ   rR   r   )
r#   rQ   r1   zKeysView[int]r2   rQ   r3   rQ   rR   rU   )Ú
__future__r   r/   r   Útypingr   Únumpyr   Úoptuna.prunersr   Úoptuna.study._study_directionr   Úoptuna.trial._stater   Úcollections.abcr   Úoptunar   r'   r5   r7   r*   r   r   ú<module>rb      s½   ðÝ "ã Û Ý  ã å %Ý 8Ý *ñ Ý(ãðØ%ðØ2@ðà
óðØ6ðàðð ðð ð	ð
 ðð óð89Ø
ð9Ø#0ð9ØBEð9ØWZð9à	ó9ô$J,�zõ J,r   