diff --git a/flaml/tune/__init__.py b/flaml/tune/__init__.py index 1418cbe4df..be26420181 100644 --- a/flaml/tune/__init__.py +++ b/flaml/tune/__init__.py @@ -35,5 +35,5 @@ ) from .sample import Categorical, Float, PolynomialExpansionSet, polynomial_expansion_set from .trial import Trial -from .tune import INCUMBENT_RESULT, report, run +from .tune import INCUMBENT_RESULT, get_run_context, report, run, use_run_context from .utils import choice diff --git a/flaml/tune/trial_runner.py b/flaml/tune/trial_runner.py index f0ec0b52c8..d7cdca526b 100644 --- a/flaml/tune/trial_runner.py +++ b/flaml/tune/trial_runner.py @@ -94,16 +94,39 @@ def process_trial_result(self, trial, result): trial.set_status(Trial.PAUSED) def stop_trial(self, trial): - """Stops trial.""" - if trial.status not in [Trial.ERROR, Trial.TERMINATED]: - if self._scheduler_alg: - self._scheduler_alg.on_trial_complete(self, trial.trial_id, trial.last_result) - self._search_alg.on_trial_complete(trial.trial_id, trial.last_result) - trial.set_status(Trial.TERMINATED) - elif self._scheduler_alg: - self._scheduler_alg.on_trial_remove(self, trial) - if trial.status == Trial.ERROR: - self._search_alg.on_trial_complete(trial.trial_id, trial.last_result, error=True) + """Stops trial. + + Holds the same per-trial lock tune.py's report() takes before its + own admission (tune.py's `_admission_lock_for`), so this cannot run + at the same instant a report is inside its locked + process_trial_result() section for the SAME trial (#996 follow-up, + sixth review point 1). Before this, a straggling background + thread's report (see the propagation machinery in tune.py: a + trainable can hand a worker thread a run context and keep reporting + through it after evaluation_function() has already returned) could + be mutating `trial.last_result`/`trial.status` and calling into the + search_alg/scheduler for this trial via process_trial_result() at + the same time run()'s own loop called stop_trial() on it directly, + unsynchronized: the two paths shared a trial but not a lock. Local + import: trial_runner is only ever imported lazily from inside + tune.run() (or after `flaml.tune`'s own package init has already + finished), specifically so tune.py has no load-time dependency on + this module; importing back here only inside the call keeps that + one-directional at import time while both places key the same lock + off the same trial object. + """ + from .tune import _admission_lock_for + + with _admission_lock_for(trial): + if trial.status not in [Trial.ERROR, Trial.TERMINATED]: + if self._scheduler_alg: + self._scheduler_alg.on_trial_complete(self, trial.trial_id, trial.last_result) + self._search_alg.on_trial_complete(trial.trial_id, trial.last_result) + trial.set_status(Trial.TERMINATED) + elif self._scheduler_alg: + self._scheduler_alg.on_trial_remove(self, trial) + if trial.status == Trial.ERROR: + self._search_alg.on_trial_complete(trial.trial_id, trial.last_result, error=True) class SequentialTrialRunner(BaseTrialRunner): diff --git a/flaml/tune/tune.py b/flaml/tune/tune.py index ea2abfe2a6..8ab972f34c 100644 --- a/flaml/tune/tune.py +++ b/flaml/tune/tune.py @@ -1,1039 +1,1759 @@ -# ! -# * Copyright (c) FLAML authors. All rights reserved. -# * Licensed under the MIT License. See LICENSE file in the -# * project root for license information. -import datetime -import os -import sys -import time -from collections import defaultdict -from typing import Callable, Dict, List, Optional, Tuple, Union - -import numpy as np - -try: - from ray import __version__ as ray_version - - assert ray_version >= "1.10.0" - from ray.tune.analysis import ExperimentAnalysis as EA -except (ImportError, AssertionError): - ray_available = False - from .analysis import ExperimentAnalysis as EA -else: - ray_available = True -import logging - -from flaml.tune.spark.utils import PySparkOvertimeMonitor, check_spark - -from .logger import logger, logger_formatter -from .result import DEFAULT_METRIC -from .trial import Trial - -try: - import mlflow -except ImportError: - mlflow = None -try: - from flaml.fabric.mlflow import MLflowIntegration, is_autolog_enabled - - internal_mlflow = True -except ImportError: - internal_mlflow = False - - -_use_ray = True -_runner = None -_verbose = 0 -_running_trial = None -_training_iteration = 0 - -INCUMBENT_RESULT = "__incumbent_result__" - - -class ExperimentAnalysis(EA): - """Class for storing the experiment results.""" - - def __init__(self, trials, metric, mode, lexico_objectives=None): - self.best_run_id = None - try: - super().__init__(self, None, trials, metric, mode) - self.lexico_objectives = lexico_objectives - except (TypeError, ValueError): - self.trials = trials - self.default_metric = metric or DEFAULT_METRIC - self.default_mode = mode - self.lexico_objectives = lexico_objectives - - @property - def best_trial(self) -> Trial: - if self.lexico_objectives is None: - return super().best_trial - else: - return self.get_best_trial(self.default_metric, self.default_mode) - - @property - def best_config(self) -> Dict: - if self.lexico_objectives is None: - return super().best_config - else: - return self.get_best_config(self.default_metric, self.default_mode) - - def lexico_best(self, trials): - results = {index: trial.last_result for index, trial in enumerate(trials) if trial.last_result} - metrics = self.lexico_objectives["metrics"] - modes = self.lexico_objectives["modes"] - f_best = {} - keys = list(results.keys()) - length = len(keys) - histories = defaultdict(list) - for time_index in range(length): - for objective, mode in zip(metrics, modes): - histories[objective].append( - results[keys[time_index]][objective] if mode == "min" else -results[keys[time_index]][objective] - ) - obj_initial = self.lexico_objectives["metrics"][0] - feasible_index = np.array([*range(len(histories[obj_initial]))]) - for k_metric, k_mode in zip(self.lexico_objectives["metrics"], self.lexico_objectives["modes"]): - k_values = np.array(histories[k_metric]) - k_target = ( - -self.lexico_objectives["targets"][k_metric] - if k_mode == "max" - else self.lexico_objectives["targets"][k_metric] - ) - feasible_value = k_values.take(feasible_index) - f_best[k_metric] = np.min(feasible_value) - - feasible_index_filter = np.where( - feasible_value - <= max( - ( - f_best[k_metric] + self.lexico_objectives["tolerances"][k_metric] - if not isinstance(self.lexico_objectives["tolerances"][k_metric], str) - else f_best[k_metric] - * (1 + 0.01 * float(self.lexico_objectives["tolerances"][k_metric].replace("%", ""))) - ), - k_target, - ) - )[0] - feasible_index = feasible_index.take(feasible_index_filter) - best_trial = trials[feasible_index[-1]] - return best_trial - - def get_best_trial( - self, - metric: Optional[str] = None, - mode: Optional[str] = None, - scope: str = "last", - filter_nan_and_inf: bool = True, - ) -> Optional[Trial]: - if self.lexico_objectives is not None: - best_trial = self.lexico_best(self.trials) - else: - best_trial = super().get_best_trial(metric, mode, scope, filter_nan_and_inf) - return best_trial - - @property - def best_result(self) -> Dict: - if self.lexico_best is None: - return super().best_result - else: - return self.best_trial.last_result - - @property - def best_iteration(self) -> List[str]: - """Help better navigate""" - best_trial = self.best_trial - best_trial_id = best_trial.trial_id - for i, trial in enumerate(self.trials): - if trial.trial_id == best_trial_id: - return i - return None - - -def report(_metric=None, **kwargs): - """A function called by the HPO application to report final or intermediate - results. - - Example: - - ```python - import time - from flaml import tune - - def compute_with_config(config): - current_time = time.time() - metric2minimize = (round(config['x'])-95000)**2 - time2eval = time.time() - current_time - tune.report(metric2minimize=metric2minimize, time2eval=time2eval) - - analysis = tune.run( - compute_with_config, - config={ - 'x': tune.lograndint(lower=1, upper=1000000), - 'y': tune.randint(lower=1, upper=1000000) - }, - metric='metric2minimize', mode='min', - num_samples=1000000, time_budget_s=60, use_ray=False) - - print(analysis.trials[-1].last_result) - ``` - - Args: - _metric: Optional default anonymous metric for ``tune.report(value)``. - (For compatibility with ray.tune.report) - **kwargs: Any key value pair to be reported. - - Raises: - StopIteration (when not using ray, i.e., _use_ray=False): - A StopIteration exception is raised if the trial has been signaled to stop. - SystemExit (when using ray): - A SystemExit exception is raised if the trial has been signaled to stop by ray. - """ - global _use_ray - global _verbose - global _running_trial - global _training_iteration - if _use_ray: - try: - from ray import __version__ as ray_version - - if ray_version.startswith("1."): - from ray import tune - - return tune.report(_metric, **kwargs) - else: # ray>=2 - from ray.air import session - - return session.report(metrics={"metric": _metric, **kwargs}) - except ImportError: - # calling tune.report() outside tune.run() - return - result = kwargs - if _metric is not None: - result[DEFAULT_METRIC] = _metric - trial = getattr(_runner, "running_trial", None) - if not trial: - return None - if _running_trial == trial: - _training_iteration += 1 - else: - _training_iteration = 0 - _running_trial = trial - result["training_iteration"] = _training_iteration - result["config"] = trial.config - if INCUMBENT_RESULT in result["config"]: - del result["config"][INCUMBENT_RESULT] - for key, value in trial.config.items(): - result["config/" + key] = value - _runner.process_trial_result(trial, result) - if _verbose > 2: - logger.info(f"result: {result}") - if trial.is_finished(): - raise StopIteration - - -def run( - evaluation_function, - config: Optional[dict] = None, - low_cost_partial_config: Optional[dict] = None, - cat_hp_cost: Optional[dict] = None, - metric: Optional[str] = None, - mode: Optional[str] = None, - time_budget_s: Union[int, float] = None, - points_to_evaluate: Optional[List[dict]] = None, - evaluated_rewards: Optional[List] = None, - resource_attr: Optional[str] = None, - min_resource: Optional[float] = None, - max_resource: Optional[float] = None, - reduction_factor: Optional[float] = None, - scheduler=None, - search_alg=None, - verbose: Optional[int] = 2, - local_dir: Optional[str] = None, - num_samples: Optional[int] = 1, - resources_per_trial: Optional[dict] = None, - config_constraints: Optional[List[Tuple[Callable[[dict], float], str, float]]] = None, - metric_constraints: Optional[List[Tuple[str, str, float]]] = None, - max_failure: Optional[int] = 100, - use_ray: Optional[bool] = False, - use_spark: Optional[bool] = False, - use_incumbent_result_in_evaluation: Optional[bool] = None, - log_file_name: Optional[str] = None, - lexico_objectives: Optional[dict] = None, - force_cancel: Optional[bool] = False, - n_concurrent_trials: Optional[int] = 0, - mlflow_exp_name: Optional[str] = None, - automl_info: Optional[Tuple[float]] = None, - extra_tag: Optional[dict] = None, - cost_attr: Optional[str] = "auto", - cost_budget: Optional[float] = None, - **ray_args, -): - """The function-based way of performing HPO. - - Example: - - ```python - import time - from flaml import tune - - def compute_with_config(config): - current_time = time.time() - metric2minimize = (round(config['x'])-95000)**2 - time2eval = time.time() - current_time - tune.report(metric2minimize=metric2minimize, time2eval=time2eval) - # if the evaluation fails unexpectedly and the exception is caught, - # and it doesn't inform the goodness of the config, - # return {} - # if the failure indicates a config is bad, - # report a bad metric value like np.inf or -np.inf - # depending on metric mode being min or max - - analysis = tune.run( - compute_with_config, - config={ - 'x': tune.lograndint(lower=1, upper=1000000), - 'y': tune.randint(lower=1, upper=1000000) - }, - metric='metric2minimize', mode='min', - num_samples=-1, time_budget_s=60, use_ray=False) - - print(analysis.trials[-1].last_result) - ``` - - Args: - evaluation_function: A user-defined evaluation function. - It takes a configuration as input, outputs a evaluation - result (can be a numerical value or a dictionary of string - and numerical value pairs) for the input configuration. - For machine learning tasks, it usually involves training and - scoring a machine learning model, e.g., through validation loss. - config: A dictionary to specify the search space. - low_cost_partial_config: A dictionary from a subset of - controlled dimensions to the initial low-cost values. - e.g., ```{'n_estimators': 4, 'max_leaves': 4}``` - - cat_hp_cost: A dictionary from a subset of categorical dimensions - to the relative cost of each choice. - e.g., ```{'tree_method': [1, 1, 2]}``` - i.e., the relative cost of the - three choices of 'tree_method' is 1, 1 and 2 respectively - metric: A string of the metric name to optimize for. - mode: A string in ['min', 'max'] to specify the objective as - minimization or maximization. - time_budget_s: int or float | The time budget in seconds. - points_to_evaluate: A list of initial hyperparameter - configurations to run first. - evaluated_rewards (list): If you have previously evaluated the - parameters passed in as points_to_evaluate you can avoid - re-running those trials by passing in the reward attributes - as a list so the optimiser can be told the results without - needing to re-compute the trial. Must be the same or shorter length than - points_to_evaluate. - e.g., - - ```python - points_to_evaluate = [ - {"b": .99, "cost_related": {"a": 3}}, - {"b": .99, "cost_related": {"a": 2}}, - ] - evaluated_rewards = [3.0] - ``` - - means that you know the reward for the first config in - points_to_evaluate is 3.0 and want to inform run(). - - resource_attr: A string to specify the resource dimension used by - the scheduler via "scheduler". - min_resource: A float of the minimal resource to use for the resource_attr. - max_resource: A float of the maximal resource to use for the resource_attr. - reduction_factor: A float of the reduction factor used for incremental - pruning. - scheduler: A scheduler for executing the experiment. Can be None, 'flaml', - 'asha' (or 'async_hyperband', 'asynchyperband') or a custom instance of the TrialScheduler class. Default is None: - in this case when resource_attr is provided, the 'flaml' scheduler will be - used, otherwise no scheduler will be used. When set 'flaml', an - authentic scheduler implemented in FLAML will be used. It does not - require users to report intermediate results in evaluation_function. - Find more details about this scheduler in this paper - https://arxiv.org/pdf/1911.04706.pdf). - When set 'asha', the input for arguments "resource_attr", - "min_resource", "max_resource" and "reduction_factor" will be passed - to ASHA's "time_attr", "max_t", "grace_period" and "reduction_factor" - respectively. You can also provide a self-defined scheduler instance - of the TrialScheduler class. When 'asha' or self-defined scheduler is - used, you usually need to report intermediate results in the evaluation - function via 'tune.report()'. - If you would like to do some cleanup opearation when the trial is stopped - by the scheduler, you can catch the `StopIteration` (when not using ray) - or `SystemExit` (when using ray) exception explicitly, - as shown in the following example. - Please find more examples using different types of schedulers - and how to set up the corresponding evaluation functions in - test/tune/test_scheduler.py, and test/tune/example_scheduler.py. - ```python - def easy_objective(config): - width, height = config["width"], config["height"] - for step in range(config["steps"]): - intermediate_score = evaluation_fn(step, width, height) - try: - tune.report(iterations=step, mean_loss=intermediate_score) - except (StopIteration, SystemExit): - # do cleanup operation here - return - ``` - search_alg: An instance/string of the search algorithm - to be used. The same instance can be used for iterative tuning. - e.g., - - ```python - from flaml import BlendSearch - algo = BlendSearch(metric='val_loss', mode='min', - space=search_space, - low_cost_partial_config=low_cost_partial_config) - for i in range(10): - analysis = tune.run(compute_with_config, - search_alg=algo, use_ray=False) - print(analysis.trials[-1].last_result) - ``` - - verbose: 0, 1, 2, or 3. If ray or spark backend is used, their verbosity will be - affected by this argument. 0 = silent, 1 = only status updates, - 2 = status and brief trial results, 3 = status and detailed trial results. - Defaults to 2. - local_dir: A string of the local dir to save ray logs if ray backend is - used; or a local dir to save the tuning log. - num_samples: An integer of the number of configs to try. Defaults to 1. - resources_per_trial: A dictionary of the hardware resources to allocate - per trial, e.g., `{'cpu': 1}`. It is only valid when using ray backend - (by setting 'use_ray = True'). It shall be used when you need to do - [parallel tuning](/docs/Use-Cases/Tune-User-Defined-Function#parallel-tuning). - config_constraints: A list of config constraints to be satisfied. - e.g., ```config_constraints = [(mem_size, '<=', 1024**3)]``` - - mem_size is a function which produces a float number for the bytes - needed for a config. - It is used to skip configs which do not fit in memory. - metric_constraints: A list of metric constraints to be satisfied. - e.g., `['precision', '>=', 0.9]`. The sign can be ">=" or "<=". - max_failure: int | the maximal consecutive number of failures to sample - a trial before the tuning is terminated. - use_ray: A boolean of whether to use ray as the backend. - use_spark: A boolean of whether to use spark as the backend. - log_file_name: A string of the log file name. Default to None. - When set to None: - if local_dir is not given, no log file is created; - if local_dir is given, the log file name will be autogenerated under local_dir. - Only valid when verbose > 0 or use_ray is True. - lexico_objectives: dict, default=None | It specifics information needed to perform multi-objective - optimization with lexicographic preferences. When lexico_objectives is not None, the arguments metric, - mode, will be invalid, and flaml's tune uses CFO - as the `search_alg`, which makes the input (if provided) `search_alg' invalid. - This dictionary shall contain the following fields of key-value pairs: - - "metrics": a list of optimization objectives with the orders reflecting the priorities/preferences of the - objectives. - - "modes" (optional): a list of optimization modes (each mode either "min" or "max") corresponding to the - objectives in the metric list. If not provided, we use "min" as the default mode for all the objectives. - - "targets" (optional): a dictionary to specify the optimization targets on the objectives. The keys are the - metric names (provided in "metric"), and the values are the numerical target values. - - "tolerances" (optional): a dictionary to specify the optimality tolerances on objectives. The keys are the metric names (provided in "metrics"), and the values are the absolute/percentage tolerance in the form of numeric/string. - E.g., - ```python - lexico_objectives = { - "metrics": ["error_rate", "pred_time"], - "modes": ["min", "min"], - "tolerances": {"error_rate": 0.01, "pred_time": 0.0}, - "targets": {"error_rate": 0.0}, - } - ``` - We also support percentage tolerance. - E.g., - ```python - lexico_objectives = { - "metrics": ["error_rate", "pred_time"], - "modes": ["min", "min"], - "tolerances": {"error_rate": "5%", "pred_time": "0%"}, - "targets": {"error_rate": 0.0}, - } - ``` - force_cancel: boolean, default=False | Whether to forcely cancel the PySpark job if overtime. - mlflow_exp_name: str, default=None | The name of the mlflow experiment. This should be specified if - enable mlflow autologging on Spark. Otherwise it will log all the results into the experiment of the - same name as the basename of main entry file. - automl_info: tuple, default=None | The information of the automl run. It should be a tuple of (mlflow_log_latency,). - n_concurrent_trials: int, default=0 | The number of concurrent trials when perform hyperparameter - tuning with Spark. Only valid when use_spark=True and spark is required: - `pip install flaml[spark]`. Please check - [here](https://spark.apache.org/docs/latest/api/python/getting_started/install.html) - for more details about installing Spark. When tune.run() is called from AutoML, it will be - overwritten by the value of `n_concurrent_trials` in AutoML. When <= 0, the concurrent trials - will be set to the number of executors. - extra_tag: dict, default=None | Extra tags to be added to the mlflow runs created by autologging. - cost_attr: None or str to specify the attribute to evaluate the cost of different trials. - Default is "auto", which means that we will automatically choose the cost attribute to use (depending - on the nature of the resource budget). When cost_attr is set to None, cost differences between different trials will be omitted - in our search algorithm. When cost_attr is set to a str different from "auto" and "time_total_s", - this cost_attr must be available in the result dict of the trial. - cost_budget: A float of the cost budget. Only valid when cost_attr is a str different from "auto" and "time_total_s". - **ray_args: keyword arguments to pass to ray.tune.run(). - Only valid when use_ray=True. - """ - global _use_ray - global _verbose - global _running_trial - global _training_iteration - global internal_mlflow - old_use_ray = _use_ray - old_verbose = _verbose - old_running_trial = _running_trial - old_training_iteration = _training_iteration - - if log_file_name: - dir_name = os.path.dirname(log_file_name) - if dir_name: - os.makedirs(dir_name, exist_ok=True) - elif local_dir and verbose > 0: - os.makedirs(local_dir, exist_ok=True) - log_file_name = os.path.join(local_dir, "tune_" + str(datetime.datetime.now()).replace(":", "-") + ".log") - if use_ray and use_spark: - raise ValueError("use_ray and use_spark cannot be both True.") - if not use_ray: - _use_ray = False - _verbose = verbose - old_handlers = logger.handlers - old_level = logger.getEffectiveLevel() - logger.handlers = [] - global _runner - old_runner = _runner - assert not ray_args, "ray_args is only valid when use_ray=True" - if ( - old_handlers - and isinstance(old_handlers[0], logging.StreamHandler) - and not isinstance(old_handlers[0], logging.FileHandler) - ): - # Add the console handler. - logger.addHandler(old_handlers[0]) - if verbose > 0: - if log_file_name: - logger.addHandler(logging.FileHandler(log_file_name)) - elif not logger.hasHandlers(): - # Add the console handler. - _ch = logging.StreamHandler(stream=sys.stdout) - _ch.setFormatter(logger_formatter) - logger.addHandler(_ch) - if verbose <= 2: - logger.setLevel(logging.INFO) - else: - logger.setLevel(logging.DEBUG) - else: - logger.setLevel(logging.CRITICAL) - - if internal_mlflow and not automl_info and (mlflow.active_run() or is_autolog_enabled()): - mlflow_integration = MLflowIntegration("tune", mlflow_exp_name, extra_tag) - evaluation_function = mlflow_integration.wrap_evaluation_function(evaluation_function) - _internal_mlflow = not automl_info # True if mlflow_integration will be used for logging - else: - _internal_mlflow = False - - from .searcher.blendsearch import CFO, BlendSearch, RandomSearch - - if lexico_objectives is not None: - if "modes" not in lexico_objectives.keys(): - lexico_objectives["modes"] = ["min"] * len(lexico_objectives["metrics"]) - for t_metric, t_mode in zip(lexico_objectives["metrics"], lexico_objectives["modes"]): - if t_metric not in lexico_objectives["tolerances"].keys(): - lexico_objectives["tolerances"][t_metric] = 0 - if t_metric not in lexico_objectives["targets"].keys(): - lexico_objectives["targets"][t_metric] = -float("inf") if t_mode == "min" else float("inf") - if search_alg is None or isinstance(search_alg, str): - if isinstance(search_alg, str): - assert search_alg in [ - "BlendSearch", - "CFO", - "CFOCat", - "RandomSearch", - ], f"search_alg={search_alg} is not recognized. 'BlendSearch', 'CFO', 'CFOcat' and 'RandomSearch' are supported." - - flaml_scheduler_resource_attr = ( - flaml_scheduler_min_resource - ) = flaml_scheduler_max_resource = flaml_scheduler_reduction_factor = None - if scheduler in (None, "flaml"): - # when scheduler is set 'flaml' or None, we will use a scheduler that is - # authentic to the search algorithms in flaml. After setting up - # the search algorithm accordingly, we need to set scheduler to - # None in case it is later used in the trial runner. - flaml_scheduler_resource_attr = resource_attr - flaml_scheduler_min_resource = min_resource - flaml_scheduler_max_resource = max_resource - flaml_scheduler_reduction_factor = reduction_factor - scheduler = None - if lexico_objectives: - # TODO: Modify after supporting BlendSearch in lexicographic optimization - SearchAlgorithm = CFO - logger.info( - f"Using search algorithm {SearchAlgorithm.__name__} for lexicographic optimization. Note that when providing other search algorithms, we use CFO instead temporarily." - ) - metric = lexico_objectives["metrics"][0] or DEFAULT_METRIC - else: - if not search_alg or search_alg == "BlendSearch": - try: - import optuna as _ - - SearchAlgorithm = BlendSearch - logger.info(f"Using search algorithm {SearchAlgorithm.__name__}.") - except ImportError: - if search_alg == "BlendSearch": - raise ValueError("To use BlendSearch, run: pip install flaml[blendsearch]") - else: - SearchAlgorithm = CFO - logger.warning("Using CFO for search. To use BlendSearch, run: pip install flaml[blendsearch]") - else: - SearchAlgorithm = locals()[search_alg] - logger.info(f"Using search algorithm {SearchAlgorithm.__name__}.") - metric = metric or DEFAULT_METRIC - search_alg = SearchAlgorithm( - metric=metric, - mode=mode, - space=config, - points_to_evaluate=points_to_evaluate, - evaluated_rewards=evaluated_rewards, - low_cost_partial_config=low_cost_partial_config, - cat_hp_cost=cat_hp_cost, - time_budget_s=time_budget_s, - num_samples=num_samples, - resource_attr=flaml_scheduler_resource_attr, - min_resource=flaml_scheduler_min_resource, - max_resource=flaml_scheduler_max_resource, - reduction_factor=flaml_scheduler_reduction_factor, - config_constraints=config_constraints, - metric_constraints=metric_constraints, - use_incumbent_result_in_evaluation=use_incumbent_result_in_evaluation, - lexico_objectives=lexico_objectives, - cost_attr=cost_attr, - cost_budget=cost_budget, - ) - else: - if metric is None or mode is None: - if lexico_objectives: - metric = lexico_objectives["metrics"][0] or metric or search_alg.metric or DEFAULT_METRIC - mode = lexico_objectives["modes"][0] or mode or search_alg.mode - else: - metric = metric or search_alg.metric or DEFAULT_METRIC - mode = mode or search_alg.mode - if ray_available and use_ray: - if ray_version.startswith("1."): - from ray.tune.suggest import ConcurrencyLimiter - else: - from ray.tune.search import ConcurrencyLimiter - else: - from flaml.tune.searcher.suggestion import ConcurrencyLimiter - if ( - search_alg.__class__.__name__ - in [ - "BlendSearch", - "CFO", - "CFOCat", - ] - and use_incumbent_result_in_evaluation is not None - ): - search_alg.use_incumbent_result_in_evaluation = use_incumbent_result_in_evaluation - searcher = search_alg.searcher if isinstance(search_alg, ConcurrencyLimiter) else search_alg - if lexico_objectives: - # TODO: Modify after supporting BlendSearch in lexicographic optimization - assert search_alg.__class__.__name__ in [ - "CFO", - ], "If lexico_objectives is not None, the search_alg must be CFO for now." - search_alg.lexico_objective = lexico_objectives - - if isinstance(searcher, BlendSearch): - setting = {} - if time_budget_s: - setting["time_budget_s"] = time_budget_s - if num_samples > 0: - setting["num_samples"] = num_samples - searcher.set_search_properties(metric, mode, config, **setting) - else: - searcher.set_search_properties(metric, mode, config) - if scheduler in ("asha", "asynchyperband", "async_hyperband"): - params = {} - # scheduler resource_dimension=resource_attr - if resource_attr: - params["time_attr"] = resource_attr - if max_resource: - params["max_t"] = max_resource - if min_resource: - params["grace_period"] = min_resource - if reduction_factor: - params["reduction_factor"] = reduction_factor - if ray_available: - from ray.tune.schedulers import ASHAScheduler - - scheduler = ASHAScheduler(**params) - if use_ray: - try: - from ray import tune - except ImportError: - raise ImportError("Failed to import ray tune. " "Please install ray[tune] or set use_ray=False") - _use_ray = True - try: - analysis = tune.run( - evaluation_function, - metric=metric, - mode=mode, - search_alg=search_alg, - scheduler=scheduler, - time_budget_s=time_budget_s, - verbose=verbose, - local_dir=local_dir, - num_samples=num_samples, - resources_per_trial=resources_per_trial, - **ray_args, - ) - if log_file_name: - with open(log_file_name, "w") as f: - for trial in analysis.trials: - f.write(f"result: {trial.last_result}\n") - return analysis - finally: - _use_ray = old_use_ray - _verbose = old_verbose - _running_trial = old_running_trial - _training_iteration = old_training_iteration - - if use_spark: - # parallel run with spark - spark_available, spark_error_msg = check_spark() - if not spark_available: - raise spark_error_msg - try: - from joblib import Parallel, delayed, parallel_backend - from joblibspark import register_spark - from pyspark.sql import SparkSession - except ImportError as e: - raise ImportError(f"{e}. Try pip install flaml[spark] or set use_spark=False.") - from flaml.tune.searcher.suggestion import ConcurrencyLimiter - - from .trial_runner import SparkTrialRunner - - register_spark() - spark = SparkSession.builder.getOrCreate() - sc = spark._jsc.sc() - num_executors = len([executor.host() for executor in sc.statusTracker().getExecutorInfos()]) - 1 - """ - By default, the number of executors is the number of VMs in the cluster. And we can - launch one trial per executor. However, sometimes we can launch more trials than - the number of executors (e.g., local mode). In this case, we can set the environment - variable `FLAML_MAX_CONCURRENT` to override the detected `num_executors`. - - `max_concurrent` is the maximum number of concurrent trials defined by `search_alg`, - `FLAML_MAX_CONCURRENT` will also be used to override `max_concurrent` if `search_alg` - is not an instance of `ConcurrencyLimiter`. - - The final number of concurrent trials is the minimum of `max_concurrent` and - `num_executors` if `n_concurrent_trials<=0` (default, automl cases), otherwise the - minimum of `max_concurrent` and `n_concurrent_trials` (tuning cases). - """ - time_start = time.time() - try: - FLAML_MAX_CONCURRENT = int(os.getenv("FLAML_MAX_CONCURRENT", 0)) - except ValueError: - FLAML_MAX_CONCURRENT = 0 - num_executors = max(num_executors, FLAML_MAX_CONCURRENT, 1) - max_spark_parallelism = max(spark.sparkContext.defaultParallelism, FLAML_MAX_CONCURRENT) - if scheduler: - scheduler.set_search_properties(metric=metric, mode=mode) - if isinstance(search_alg, ConcurrencyLimiter): - max_concurrent = max(1, search_alg.max_concurrent) - else: - max_concurrent = max(1, max_spark_parallelism) - passed_in_n_concurrent_trials = max(n_concurrent_trials, max_concurrent) - n_concurrent_trials = min( - n_concurrent_trials if n_concurrent_trials > 0 else num_executors, - max_concurrent, - ) - if n_concurrent_trials < passed_in_n_concurrent_trials: - logger.warning( - f"The actual concurrent trials is {n_concurrent_trials}. You can set the environment " - f"variable `FLAML_MAX_CONCURRENT` to '{passed_in_n_concurrent_trials}' to override the detected num of executors." - ) - with parallel_backend("spark"): - with Parallel(n_jobs=n_concurrent_trials, verbose=max(0, (verbose - 1) * 50)) as parallel: - try: - _runner = SparkTrialRunner( - search_alg=search_alg, - scheduler=scheduler, - metric=metric, - mode=mode, - ) - num_trials = 0 - if time_budget_s is None: - time_budget_s = np.inf - num_failures = 0 - upperbound_num_failures = (len(evaluated_rewards) if evaluated_rewards else 0) + max_failure - logger.debug(f"automl_info: {automl_info}") - while ( - time.time() - time_start < time_budget_s - and (num_samples < 0 or num_trials < num_samples) - and num_failures < upperbound_num_failures - ): - if automl_info and automl_info[1] == "all" and automl_info[0] > 0 and time_budget_s < np.inf: - time_budget_s -= automl_info[0] * n_concurrent_trials - logger.debug(f"Remaining time budget with mlflow log latency: {time_budget_s} seconds.") - while len(_runner.running_trials) < n_concurrent_trials: - # suggest trials for spark - trial_next = _runner.step() - if trial_next: - num_trials += 1 - else: - num_failures += 1 # break with upperbound_num_failures consecutive failures - logger.debug(f"consecutive failures is {num_failures}") - if num_failures >= upperbound_num_failures: - break - trials_to_run = _runner.running_trials - if not trials_to_run: - logger.warning(f"fail to sample a trial for {max_failure} times in a row, stopping.") - break - logger.info( - f"Number of trials: {num_trials}/{num_samples}, {len(_runner.running_trials)} RUNNING," - f" {len(_runner._trials) - len(_runner.running_trials)} TERMINATED" - ) - logger.debug( - f"Configs of Trials to run: {[trial_to_run.config for trial_to_run in trials_to_run]}" - ) - results = None - with PySparkOvertimeMonitor(time_start, time_budget_s, force_cancel, parallel=parallel): - try: - results = parallel( - delayed(evaluation_function)(trial_to_run.config) for trial_to_run in trials_to_run - ) - except RuntimeError as e: - logger.warning(f"RuntimeError: {e}") - results = None - logger.info( - "Encountered RuntimeError. Waiting 10 seconds for Spark cluster to recover before retrying." - ) - time.sleep(10) - # results = [evaluation_function(trial_to_run.config) for trial_to_run in trials_to_run] - while results: - result = results.pop(0) - trial_to_run = trials_to_run[0] - _runner.running_trial = trial_to_run - if result is not None: - if _internal_mlflow: - mlflow_integration.record_trial(result, trial_to_run, metric) - - if isinstance(result, dict): - if result: - logger.info(f"Brief result: {result}") - report(**result) - else: - # When the result returned is an empty dict, set the trial status to error - trial_to_run.set_status(Trial.ERROR) - else: - logger.info("Brief result: {metric: result}") - report(_metric=result) - _runner.stop_trial(trial_to_run) - num_failures = 0 - analysis = ExperimentAnalysis( - _runner.get_trials(), - metric=metric, - mode=mode, - lexico_objectives=lexico_objectives, - ) - analysis.search_space = config - - if _internal_mlflow: - mlflow_integration.log_tune(analysis, metric) - # try: - # _best_config = analysis.best_config - # except Exception: - # _best_config = None - # if _best_config: - # parallel( - # delayed(mlflow_integration.retrain)(evaluation_function, analysis.best_config) - # for dummy in [0] - # ) - - return analysis - finally: - # recover the global variables in case of nested run - _use_ray = old_use_ray - _verbose = old_verbose - _running_trial = old_running_trial - _training_iteration = old_training_iteration - if not use_ray: - _runner = old_runner - logger.handlers = old_handlers - logger.setLevel(old_level) - if _internal_mlflow: - mlflow_integration.adopt_children() - - # simple sequential run without using tune.run() from ray - time_start = time.time() - _use_ray = False - if scheduler: - scheduler.set_search_properties(metric=metric, mode=mode) - from .trial_runner import SequentialTrialRunner - - try: - _runner = SequentialTrialRunner( - search_alg=search_alg, - scheduler=scheduler, - metric=metric, - mode=mode, - ) - num_trials = 0 - if time_budget_s is None: - time_budget_s = np.inf - num_failures = 0 - upperbound_num_failures = (len(evaluated_rewards) if evaluated_rewards else 0) + max_failure - while ( - time.time() - time_start < time_budget_s - and (num_samples < 0 or num_trials < num_samples) - and num_failures < upperbound_num_failures - ): - trial_to_run = _runner.step() - if trial_to_run: - num_trials += 1 - if verbose: - logger.info(f"trial {num_trials} config: {trial_to_run.config}") - result = None - with PySparkOvertimeMonitor(time_start, time_budget_s, force_cancel): - result = evaluation_function(trial_to_run.config) - logger.debug(f"result in tune: {trial_to_run}, {result}") - if result is not None: - if _internal_mlflow: - mlflow_integration.record_trial(result, trial_to_run, metric) - - if isinstance(result, dict): - if result: - report(**result) - else: - # When the result returned is an empty dict, set the trial status to error - trial_to_run.set_status(Trial.ERROR) - else: - report(_metric=result) - _runner.stop_trial(trial_to_run) - num_failures = 0 - if trial_to_run.last_result is None: - # application stops tuning by returning None - # TODO document this feature when it is finalized - break - else: - # break with upperbound_num_failures consecutive failures - num_failures += 1 - if num_failures == upperbound_num_failures: - logger.warning(f"fail to sample a trial for {max_failure} times in a row, stopping.") - analysis = ExperimentAnalysis( - _runner.get_trials(), - metric=metric, - mode=mode, - lexico_objectives=lexico_objectives, - ) - analysis.search_space = config - if _internal_mlflow: - mlflow_integration.log_tune(analysis, metric) - if analysis.best_run_id is not None: - logger.info(f"Best MLflow run name: {analysis.best_run_name}") - logger.info(f"Best MLflow run id: {analysis.best_run_id}") - # try: - # _best_config = analysis.best_config - # except Exception: - # _best_config = None - # if _best_config: - # mlflow_integration.retrain(evaluation_function, analysis.best_config) - - return analysis - finally: - # recover the global variables in case of nested run - _use_ray = old_use_ray - _verbose = old_verbose - _running_trial = old_running_trial - _training_iteration = old_training_iteration - if not use_ray: - _runner = old_runner - logger.handlers = old_handlers - logger.setLevel(old_level) - if _internal_mlflow: - mlflow_integration.adopt_children() - - -class Tuner: - """Tuner is the class-based way of launching hyperparameter tuning jobs compatible with Ray Tune 2. - - Args: - trainable: A user-defined evaluation function. - It takes a configuration as input, outputs a evaluation - result (can be a numerical value or a dictionary of string - and numerical value pairs) for the input configuration. - For machine learning tasks, it usually involves training and - scoring a machine learning model, e.g., through validation loss. - param_space: Search space of the tuning job. - One thing to note is that both preprocessor and dataset can be tuned here. - tune_config: Tuning algorithm specific configs. - Refer to ray.tune.tune_config.TuneConfig for more info. - run_config: Runtime configuration that is specific to individual trials. - If passed, this will overwrite the run config passed to the Trainer, - if applicable. Refer to ray.air.config.RunConfig for more info. - - Usage pattern: - - .. code-block:: python - - from sklearn.datasets import load_breast_cancer - - from ray import tune - from ray.data import from_pandas - from ray.air.config import RunConfig, ScalingConfig - from ray.train.xgboost import XGBoostTrainer - from ray.tune.tuner import Tuner - - def get_dataset(): - data_raw = load_breast_cancer(as_frame=True) - dataset_df = data_raw["data"] - dataset_df["target"] = data_raw["target"] - dataset = from_pandas(dataset_df) - return dataset - - trainer = XGBoostTrainer( - label_column="target", - params={}, - datasets={"train": get_dataset()}, - ) - - param_space = { - "scaling_config": ScalingConfig( - num_workers=tune.grid_search([2, 4]), - resources_per_worker={ - "CPU": tune.grid_search([1, 2]), - }, - ), - # You can even grid search various datasets in Tune. - # "datasets": { - # "train": tune.grid_search( - # [ds1, ds2] - # ), - # }, - "params": { - "objective": "binary:logistic", - "tree_method": "approx", - "eval_metric": ["logloss", "error"], - "eta": tune.loguniform(1e-4, 1e-1), - "subsample": tune.uniform(0.5, 1.0), - "max_depth": tune.randint(1, 9), - }, - } - tuner = Tuner(trainable=trainer, param_space=param_space, - run_config=RunConfig(name="my_tune_run")) - analysis = tuner.fit() - - To retry a failed tune run, you can then do - - .. code-block:: python - - tuner = Tuner.restore(experiment_checkpoint_dir) - tuner.fit() - - ``experiment_checkpoint_dir`` can be easily located near the end of the - console output of your first failed run. - """ +# ! +# * Copyright (c) FLAML authors. All rights reserved. +# * Licensed under the MIT License. See LICENSE file in the +# * project root for license information. +import concurrent.futures +import contextlib +import contextvars +import datetime +import os +import sys +import threading +import time +import weakref +from collections import defaultdict +from typing import Callable, Dict, List, Optional, Tuple, Union + +import numpy as np + +try: + from ray import __version__ as ray_version + + assert ray_version >= "1.10.0" + from ray.tune.analysis import ExperimentAnalysis as EA +except (ImportError, AssertionError): + ray_available = False + from .analysis import ExperimentAnalysis as EA +else: + ray_available = True +import logging + +from flaml.tune.spark.utils import PySparkOvertimeMonitor, check_spark + +from .logger import logger, logger_formatter +from .result import DEFAULT_METRIC +from .trial import Trial + +try: + import mlflow +except ImportError: + mlflow = None +try: + from flaml.fabric.mlflow import MLflowIntegration, is_autolog_enabled + + internal_mlflow = True +except ImportError: + internal_mlflow = False + + +class _TuneState(threading.local): + """Per-thread run() state. + + report()/run() coordinate through this state (use_ray, runner, verbose, + log_run_id) on the thread that is actually driving a tune.run() call. A + plain module global here is shared by every thread, so two threads + calling tune.run() concurrently overwrite each other's runner/trial + bookkeeping mid-flight (#996). threading.local's __init__ re-runs on + each thread's first access, so every thread starts from these same + defaults without an explicit per-thread init call. + + Being per-thread means a thread that a trainable spawns on its own (a + worker/callback thread that later calls tune.report()) starts from + these defaults too, with runner=None: report() then falls back to + _propagated_context below rather than this thread's own (empty) state. + running_trial/training_iteration used to live here too; see + _RunContext and _next_training_iteration for where that bookkeeping + went and why (#996 follow-up, second review points 1 and 2). + """ + + def __init__(self): + self.use_ray = True + self.runner = None + self.verbose = 0 + self.log_run_id = None + + +_state = _TuneState() + +INCUMBENT_RESULT = "__incumbent_result__" + + +class _RunScopedFilter(logging.Filter): + """Passes only log records emitted while the emitting thread is "inside" + the run() call that owns this filter. + + Concurrent/nested tune.run() calls each add their own Handler to the + shared `flaml.tune.logger` logger instead of replacing `logger.handlers` + wholesale, so one run's handler is never wiped out by another's setup or + teardown (#996 follow-up: the shared logger was still corrupted the same + way the five state globals used to be). This filter is what keeps one + run's records out of another's handler: matching is against thread-local + `_state.log_run_id`, so a nested run on the SAME thread resolves to + "innermost run wins" (the same behavior the old module-global design + had), while two DIFFERENT threads never see each other's records at all, + since `_state` is thread-local and Python's logging dispatch runs + synchronously on the emitting thread. + + An earlier version of this filter let a record carrying + `record.flaml_tune_unscoped = True` bypass the match, reaching every + active run's handler regardless of which thread emitted it (#996 + follow-up, seventh review), for the one caller that has no run of its + own to match against: the context-less-report warning below. That + bypass is gone (#996 follow-up, eighth review point 3): it put an + unattributed diagnostic into every OTHER concurrently active run's log + file too, which is a different, self-inflicted instance of the same + misattribution class this whole filter exists to prevent. The warning + now goes through `_diagnostics_logger` below instead, a logger this + filter is never attached to, so it needs no bypass here at all. + """ + + def __init__(self, run_id): + super().__init__() + self._run_id = run_id + + def filter(self, record): + return _state.log_run_id is self._run_id + + +# Bookkeeping for the shared logger's OWN level (logger.setLevel), which is a +# single process-global value distinct from any one run's Handler.setLevel(). +# Each concurrently active run() contributes its desired level; the logger is +# kept at the most permissive (lowest) of those so no run's handler is +# starved by another run's stricter request, and it is restored to its +# pre-any-run value once the last active run exits. _logger_state_lock only +# ever guards this short bookkeeping section, never a run's full duration. +_logger_state_lock = threading.Lock() +_active_log_levels: List[int] = [] +_logger_pristine_level: Optional[int] = None + + +def _logger_level_enter(level: int) -> None: + """Raise (numerically lower) the shared logger's level for the duration + of at least one active run, remembering the level to restore it to once + none are left. + + Saves `logger.level`, the logger's OWN unresolved level (0/NOTSET when + the logger has never had `setLevel()` called on it and is inheriting + from its ancestors), not `logger.getEffectiveLevel()` (#996 follow-up, + sixth review point 4). The two read the same value only by coincidence, + whenever the logger already had an explicit level of its own; a fresh + `flaml.tune` logger has none, so `getEffectiveLevel()` resolves up to + whatever the root logger happens to be at (commonly WARNING) and + `_logger_level_exit()` used to bake that RESOLVED number back in as an + explicit `logger.level` via `setLevel()`, which is a different, and + permanent, configuration: the logger would no longer follow a later + change to the root logger's level the way NOTSET inheritance does. + `logger.level` is exactly the value `setLevel()` needs to reproduce the + original state, inheritance included. + """ + global _logger_pristine_level + with _logger_state_lock: + if not _active_log_levels: + _logger_pristine_level = logger.level + _active_log_levels.append(level) + logger.setLevel(min(_active_log_levels)) + + +def _logger_level_exit(level: int) -> None: + with _logger_state_lock: + _active_log_levels.remove(level) + logger.setLevel(min(_active_log_levels) if _active_log_levels else _logger_pristine_level) + + +class _RunContext: + """Snapshot of one tune.run() call's active state (#996 follow-up). + + Invariant: there is one active run-context per concurrent tune.run() + call, and it must be visible to that run's own worker/callback threads, + and invisible to any other concurrent run()'s threads. + + `tune.report()` reports against `_state.runner`, which is thread-local. + If a trainable spawns its own worker/callback thread and that thread + calls `tune.report()`, the worker thread has no runner attached (its + `_TuneState` just initialized to defaults) and the report used to be + silently dropped, same as calling tune.report() outside of tune.run() + entirely (#996 follow-up, second review point 1). run() now sets + _propagated_context (below) around every evaluation_function() call it + makes on its own driving thread, and a plain threading.Thread started + from inside that call automatically inherits it, see + _install_thread_context_propagation(). get_run_context()/ + use_run_context() remain the explicit escape hatch for handoffs that + patch can't reach (a persistent thread-pool executor whose worker + threads outlive any single submitted task, for instance). + + training_iteration does NOT live here (it used to, see + _next_training_iteration for why that broke synchronization). + + log_run_id (#996 follow-up, third review point 4) is the same value + _RunScopedFilter matches against `_state.log_run_id`; carrying it here + means the one propagation mechanism below (thread or executor) also + covers "this worker's log records belong in the run log", not just + report() routing. + """ + + __slots__ = ("use_ray", "runner", "verbose", "running_trial", "log_run_id") + + def __init__(self, use_ray, runner, verbose, running_trial, log_run_id=None): + self.use_ray = use_ray + self.runner = runner + self.verbose = verbose + self.running_trial = running_trial + self.log_run_id = log_run_id + + +# Ambient propagation channel report() consults when its own thread's +# _state.runner is None (#996 follow-up, second review point 1). Set by +# use_run_context() for the duration of its `with` block, and by run() +# around each evaluation_function() call on the sequential (non-ray, +# non-spark) path. A bare threading.Thread does NOT inherit a +# contextvars.ContextVar value the way an asyncio Task does: CPython +# gives every new OS thread its own empty top-level Context, so this by +# itself only reaches use_run_context() callers, not a worker thread a +# trainable spawns on its own with no FLAML-specific code. Pairing it with +# _install_thread_context_propagation() and +# _install_executor_context_propagation() below is what makes that second, +# more common case ("existing trainables ... unless callers adopt the new +# context API", per review) work with no trainable-side change. +_propagated_context: "contextvars.ContextVar[Optional[_RunContext]]" = contextvars.ContextVar( + "flaml_tune_propagated_context", default=None +) + + +def _install_thread_context_propagation() -> None: + """Make a plain threading.Thread inherit the calling thread's active + _RunContext, process-wide, once. + + Without this, `_propagated_context` set on the thread driving + tune.run() is invisible to a `threading.Thread(...)` a trainable spawns + from inside its own evaluation_function. contextvars are per-OS-thread + in CPython by default, same as threading.local, and only asyncio Task + creation (or an explicit Context.run()) copies the parent's bindings. + + Two changes from the first version of this patch (#996 follow-up, third + review): + + - Captures and restores only `_propagated_context`'s own _RunContext + value (plus setting `_state.log_run_id` from it), never + `contextvars.copy_context()`. The full-context copy dragged every + ambient application ContextVar into every new thread process-wide, + including ones with nothing to do with tuning (third review point + 3); this only ever touches the one FLAML-owned value. + - Wraps `Thread._bootstrap_inner` instead of `Thread.run`. + `_bootstrap_inner` is what CPython's Thread._bootstrap() actually + calls, and it invokes `self.run()` internally regardless of which + `run()` that resolves to, so a Thread SUBCLASS overriding `run()` + (idiomatic and common) is still covered; the original patch replaced + only the base class's `run` attribute, which a subclass's own `run` + shadows and the patch then never runs (third review point 2). + + A thread started outside any tune.run() call captures no context and + behaves exactly as before. Idempotent: a second import/call is a no-op. + Does not help a persistent thread pool's already-running worker + threads; see _install_executor_context_propagation() for that case. + """ + if getattr(threading.Thread, "_flaml_tune_context_propagation", False): + return + _orig_start = threading.Thread.start + _orig_bootstrap_inner = threading.Thread._bootstrap_inner + + def _start(self, *args, **kwargs): + self._flaml_tune_ctx = _propagated_context.get() + return _orig_start(self, *args, **kwargs) + + def _bootstrap_inner(self, *args, **kwargs): + ctx = getattr(self, "_flaml_tune_ctx", None) + if ctx is None: + return _orig_bootstrap_inner(self, *args, **kwargs) + token = _propagated_context.set(ctx) + _state.log_run_id = ctx.log_run_id + try: + return _orig_bootstrap_inner(self, *args, **kwargs) + finally: + _propagated_context.reset(token) + + threading.Thread.start = _start + threading.Thread._bootstrap_inner = _bootstrap_inner + threading.Thread._flaml_tune_context_propagation = True + + +def _install_executor_context_propagation() -> None: + """Make concurrent.futures.ThreadPoolExecutor.submit() propagate the + submitting thread's active _RunContext to the task it submits (#996 + follow-up, third review point 1, "pre-warmed and cross-trial executor + reuse"). + + A pool's worker threads call Thread.start() once, when the pool spins + them up, not once per submitted task, so _install_thread_context_ + propagation()'s capture-at-start() never sees a context for a task + submitted to an already-running worker: that worker's start() ran (if + at all) before this task's context existed. Verified directly: + ThreadPoolExecutor.submit() does not propagate a plain + contextvars.ContextVar to an already-warmed-up worker either (CPython + does not give submit() the automatic inheritance asyncio Task creation + gets), so this is not reachable by patching Thread at any capture + point; submit() itself has to be the interception point, applied once + per task rather than once per worker thread. + + Wraps the submitted callable rather than the worker thread: the same + worker thread runs many tasks across its lifetime, each potentially + from a different tune.run() call (or none), so the context has to be + attached and detached per task, not once for the thread. Idempotent: + a second import/call is a no-op. map() is covered for free, since + concurrent.futures.Executor.map() calls self.submit() per item. + """ + if getattr(concurrent.futures.ThreadPoolExecutor, "_flaml_tune_context_propagation", False): + return + _orig_submit = concurrent.futures.ThreadPoolExecutor.submit + + def _submit(self, fn, *args, **kwargs): + ctx = _propagated_context.get() + if ctx is None: + return _orig_submit(self, fn, *args, **kwargs) + + def _flaml_tune_wrapped(*a, **kw): + token = _propagated_context.set(ctx) + prior_log_run_id = _state.log_run_id + _state.log_run_id = ctx.log_run_id + try: + return fn(*a, **kw) + finally: + _propagated_context.reset(token) + _state.log_run_id = prior_log_run_id + + return _orig_submit(self, _flaml_tune_wrapped, *args, **kwargs) + + concurrent.futures.ThreadPoolExecutor.submit = _submit + concurrent.futures.ThreadPoolExecutor._flaml_tune_context_propagation = True + + +_install_thread_context_propagation() +_install_executor_context_propagation() + +# Per-trial training_iteration bookkeeping (#996 follow-up, second review +# point 2). See _next_training_iteration for the invariant this maintains. +_trial_iteration_lock = threading.Lock() +_trial_iteration: "weakref.WeakKeyDictionary" = weakref.WeakKeyDictionary() + +# Per-trial admission lock (#996 follow-up, fifth review point 2). See +# _admission_lock_for for the invariant this maintains. A separate lock from +# _trial_iteration_lock above: _next_training_iteration() is called from +# inside report() while this lock may already be held, and reusing one +# process-wide lock for both would self-deadlock a non-reentrant +# threading.Lock the first time a report actually reaches that call. +_admission_locks_guard = threading.Lock() +_admission_locks: "weakref.WeakKeyDictionary" = weakref.WeakKeyDictionary() + + +def _admission_lock_for(trial) -> threading.Lock: + """Return the one lock guarding admission of a report into `trial`. + + report()'s is_finished() check (whether to bother building a result at + all) and its actual admission (runner.process_trial_result(), which is + what can transition the trial to finished) used to be two unsynchronized + steps: a second, concurrent report() for the SAME trial could read + is_finished() as False, then have the FIRST report's scheduler decision + finish the trial before the second one reached process_trial_result(), + which still wrote through and silently replaced the trial's real final + result with the stale one (#996 follow-up, fifth review point 2). This + lock makes the re-check immediately before process_trial_result() + atomic with the write, so a report that loses that race is dropped + instead of admitted. + + Keyed per trial (a WeakKeyDictionary, same pattern as + _trial_iteration above) rather than one global lock, so reports for + different trials never serialize against each other. + """ + with _admission_locks_guard: + lock = _admission_locks.get(trial) + if lock is None: + lock = threading.Lock() + _admission_locks[trial] = lock + return lock + + +# Every currently active SEQUENTIAL (non-ray) run's runner, keyed to its own +# log_run_id/verbose (#996 follow-up, eighth review point 1). A +# WeakKeyDictionary: an entry is dropped once its run() call restores +# _state.runner and nothing else references the old runner, so this needs no +# separate cleanup beyond the explicit .pop() run() does on the way out +# (belt-and-suspenders against a runner outliving its run() call some other +# way). Guarded by its own lock, never held across a report() call's own +# work, only the lookup. +_active_runners_lock = threading.Lock() +_active_runners: "weakref.WeakKeyDictionary" = weakref.WeakKeyDictionary() + + +def _resolve_unambiguous_active_runner(): + """Return the sole active sequential run's (runner, log_run_id, verbose), + or (None, None, None) if none are active. + + Only called for a context-less report() (#996 follow-up, eighth review + point 1): no thread-local runner and no propagated/explicit _RunContext, + the shape a generic queue/callback worker started before tune.run() + exists has, unchanged, and never calling get_run_context()/ + use_run_context() itself. Automatic per-dispatch propagation for such a + worker is still not being added (SequentialTrialRunner.step() reassigns + running_trial every step, so a live-resolved trial can still be the + WRONG one if this same run has already moved past the trial the report + was actually produced for by the time it is handled (see report()'s + own docstring). What changed: when exactly one sequential tune.run() is + active anywhere in the process, a context-less report can only mean + that one run, which is exactly the fallback the pre-#996 module-global + design gave every caller unconditionally (report() read whichever + runner the one shared global held, live, with the same delayed-item + risk this still carries within that single run). Restoring it for this + one unambiguous case preserves the working common shape (a single + tune.run() call, one worker draining a queue produced during it) instead + of dropping it outright. + + Raises RuntimeError when two or more sequential runs are active at + once: which one a context-less report belongs to is then genuinely + undecidable (there is no signal on the reporting thread that says + which), and silently guessing would attribute one run's metric to a + different run's trial, worse than dropping it. Callers who need this + to work under real concurrent runs must capture get_run_context() where + the work is produced and attach it with use_run_context() where it is + processed; that path already works today and is unaffected by this. + """ + with _active_runners_lock: + candidates = list(_active_runners.items()) + if not candidates: + return None, None, None + if len(candidates) > 1: + raise RuntimeError( + "tune.report() was called from a thread with no active run attached to it " + "(not the thread driving tune.run(), and no context was propagated or " + "explicitly attached), and more than one tune.run() call is concurrently " + "active in this process right now, so which run this report belongs to " + "cannot be determined safely. Capture get_run_context() in the code that " + "queues the work and attach it with use_run_context() where the item is " + "processed." + ) + runner, (log_run_id, verbose) = candidates[0] + return runner, log_run_id, verbose + + +def _next_training_iteration(trial) -> int: + """Return trial's next training_iteration, as a counter shared by every + thread that reports for this trial, whichever thread that is. + + training_iteration used to be a plain int on _RunContext: a snapshot + get_run_context() took of _state.training_iteration, copied onto the + receiving thread's own _state by use_run_context(), and discarded (via + _state restore) when that thread's `with` block exited. The driving + thread's own copy was never updated by a worker thread's reports, so + every get_run_context() call handed out the same stale snapshot the + driving thread's own copy still held, 0 if the driving thread never + reports directly itself, which is the common case when a worker thread + does the reporting instead. Every one of those propagated reports then + repeated the SAME training_iteration (verified: 0, for three separate + handoffs in a row), instead of the trial's count going up, which is + what a scheduler/searcher that orders trials by training_iteration + (ASHA and similar) needs to see. + + The trial object itself is the one thing every one of those threads + already holds a reference to in common (_state.runner.running_trial on + the driving thread, ctx.running_trial on a propagated one), so keying + the counter on the trial, instead of copying it through whichever + thread or context happens to be reporting, makes it actually shared. + The lock makes the read-increment-write atomic across threads; the + WeakKeyDictionary drops a trial's entry once nothing else references + it, so finished trials need no separate cleanup. + """ + with _trial_iteration_lock: + iteration = _trial_iteration.get(trial, -1) + 1 + _trial_iteration[trial] = iteration + return iteration + + +def get_run_context() -> Optional["_RunContext"]: + """Snapshot the calling thread's active tune.run() state. + + Call this from the thread tune.run() is driving the trainable on (e.g. + at the top of the trainable, before spawning a helper thread). Returns + None if this thread is not currently inside a tune.run() call, in which + case there is nothing to propagate. + + Existing trainables do not need to call this: a plain worker thread, or + a task submitted to a ThreadPoolExecutor, already gets a context + automatically, via _propagated_context and + _install_thread_context_propagation()/ + _install_executor_context_propagation(). This (and use_run_context()) + stay as the explicit form for the cases neither patch can reach (a + Thread subclass that also overrides _bootstrap_inner itself, for + instance, or a non-stdlib executor). + """ + if _state.runner is None: + return None + return _RunContext(_state.use_ray, _state.runner, _state.verbose, _state.runner.running_trial, _state.log_run_id) + + +@contextlib.contextmanager +def use_run_context(ctx: Optional["_RunContext"]): + """Attach a context captured by get_run_context() to the calling thread. + + `tune.report()` calls made inside the `with` block report against the + run `ctx` was captured from, as if they were made on the driving thread. + `ctx=None` is accepted and is a no-op, so callers do not need to special + case "this thread never got a context". + + Also attaches ctx.log_run_id to this thread's _state for the duration, + so log records emitted inside the `with` block land in the owning run's + log the same way a report() call routes to the owning trial (#996 + follow-up, third review point 4). + + Note this only propagates report() bookkeeping (which runner, which + trial) and log routing; it does not add any locking around the shared + TrialRunner, so this is meant for a single worker thread computing a + result and handing it off (report(), then join), not for multiple + threads reporting against the same trial truly concurrently. + """ + if ctx is None: + yield + return + token = _propagated_context.set(ctx) + prior_log_run_id = _state.log_run_id + _state.log_run_id = ctx.log_run_id + try: + yield + finally: + _propagated_context.reset(token) + _state.log_run_id = prior_log_run_id + + +class ExperimentAnalysis(EA): + """Class for storing the experiment results.""" + + def __init__(self, trials, metric, mode, lexico_objectives=None): + self.best_run_id = None + try: + super().__init__(self, None, trials, metric, mode) + self.lexico_objectives = lexico_objectives + except (TypeError, ValueError): + self.trials = trials + self.default_metric = metric or DEFAULT_METRIC + self.default_mode = mode + self.lexico_objectives = lexico_objectives + + @property + def best_trial(self) -> Trial: + if self.lexico_objectives is None: + return super().best_trial + else: + return self.get_best_trial(self.default_metric, self.default_mode) + + @property + def best_config(self) -> Dict: + if self.lexico_objectives is None: + return super().best_config + else: + return self.get_best_config(self.default_metric, self.default_mode) + + def lexico_best(self, trials): + results = {index: trial.last_result for index, trial in enumerate(trials) if trial.last_result} + metrics = self.lexico_objectives["metrics"] + modes = self.lexico_objectives["modes"] + f_best = {} + keys = list(results.keys()) + length = len(keys) + histories = defaultdict(list) + for time_index in range(length): + for objective, mode in zip(metrics, modes): + histories[objective].append( + results[keys[time_index]][objective] if mode == "min" else -results[keys[time_index]][objective] + ) + obj_initial = self.lexico_objectives["metrics"][0] + feasible_index = np.array([*range(len(histories[obj_initial]))]) + for k_metric, k_mode in zip(self.lexico_objectives["metrics"], self.lexico_objectives["modes"]): + k_values = np.array(histories[k_metric]) + k_target = ( + -self.lexico_objectives["targets"][k_metric] + if k_mode == "max" + else self.lexico_objectives["targets"][k_metric] + ) + feasible_value = k_values.take(feasible_index) + f_best[k_metric] = np.min(feasible_value) + + feasible_index_filter = np.where( + feasible_value + <= max( + ( + f_best[k_metric] + self.lexico_objectives["tolerances"][k_metric] + if not isinstance(self.lexico_objectives["tolerances"][k_metric], str) + else f_best[k_metric] + * (1 + 0.01 * float(self.lexico_objectives["tolerances"][k_metric].replace("%", ""))) + ), + k_target, + ) + )[0] + feasible_index = feasible_index.take(feasible_index_filter) + best_trial = trials[feasible_index[-1]] + return best_trial + + def get_best_trial( + self, + metric: Optional[str] = None, + mode: Optional[str] = None, + scope: str = "last", + filter_nan_and_inf: bool = True, + ) -> Optional[Trial]: + if self.lexico_objectives is not None: + best_trial = self.lexico_best(self.trials) + else: + best_trial = super().get_best_trial(metric, mode, scope, filter_nan_and_inf) + return best_trial + + @property + def best_result(self) -> Dict: + if self.lexico_best is None: + return super().best_result + else: + return self.best_trial.last_result + + @property + def best_iteration(self) -> List[str]: + """Help better navigate""" + best_trial = self.best_trial + best_trial_id = best_trial.trial_id + for i, trial in enumerate(self.trials): + if trial.trial_id == best_trial_id: + return i + return None + + +# Dedicated to a diagnostic that cannot be attributed to any one active run +# (#996 follow-up, eighth review point 3): a distinct logger, in its own +# branch of the hierarchy ("flaml.tune.diagnostics", a sibling of +# "flaml.tune.logger"/`logger` below, never "flaml.tune.logger" itself), so +# run() never attaches a run-owned FileHandler/StreamHandler to it and an +# unattributed record is never written into one of those files, which the +# `flaml_tune_unscoped` bypass this replaced did exactly (see +# _RunScopedFilter above). Left at its default `propagate=True`, unlike +# `logger`, so a record still reaches somewhere: up to "flaml.tune", then +# "flaml" (INFO by default, flaml/__init__.py), then root, landing on +# whatever handler the caller's own logging config has there, or on +# Python's `logging.lastResort` (stderr) if none. Verified directly: a +# handler with no run-scoped filter on it at all would have worked too, but +# would need run() to manage its lifecycle the same way it manages +# run-owned handlers, which is the coupling this is meant to avoid. +_diagnostics_logger = logging.getLogger("flaml.tune.diagnostics") + +# Guards the warning below (#996 follow-up, seventh and eighth review). A +# generic queue/callback worker started before any tune.run() call exists +# captures no _propagated_context at Thread.start() time (that mechanism +# only reaches a thread started AFTER a run is already active) and never +# calls get_run_context()/use_run_context() itself, so a report() it makes +# while no run is active anywhere in the process, or while more than one +# is (see _resolve_unambiguous_active_runner, which raises for that second +# case instead of reaching here), lands on the "no trial to attribute this +# to" branch below. That branch can run once per dequeued item on a hot +# worker loop, so this is scoped to fire at most once per call to run(), +# not once per report(): run() resets the flag at the top of every call +# (#996 follow-up, eighth review point 2), so a LATER run that hits the +# same unsupported shape warns again instead of staying silent forever +# after the first affected run in the process. A nested tune.run() call on +# the same thread also resets it, which can make an outer run's own +# warning fire twice across one process lifetime; that is the direction to +# fail in, not the reverse. +_context_less_report_warned = False +_context_less_report_warned_lock = threading.Lock() + + +def _warn_context_less_report_once() -> None: + global _context_less_report_warned + if _context_less_report_warned: + return + with _context_less_report_warned_lock: + if _context_less_report_warned: + return + _context_less_report_warned = True + _diagnostics_logger.warning( + "tune.report() was called from a thread with no active run to attribute it " + "to (not the thread driving tune.run(), and no context was propagated or " + "explicitly attached), and no tune.run() call is active in this process " + "right now either, so there is nothing to attribute it to even by " + "inference. The report is dropped, not recorded against any trial, and " + "this will not be logged again for this run. A worker thread started " + "before tune.run() is called, such as a persistent queue consumer, whose " + "items are processed while exactly one tune.run() call is active, is " + "handled automatically; this warning means that was not the case here. " + "Capture get_run_context() in the code that queues the work and attach it " + "with use_run_context() where the item is processed.", + ) + + +def report(_metric=None, **kwargs): + """A function called by the HPO application to report final or intermediate + results. + + Example: + + ```python + import time + from flaml import tune + + def compute_with_config(config): + current_time = time.time() + metric2minimize = (round(config['x'])-95000)**2 + time2eval = time.time() - current_time + tune.report(metric2minimize=metric2minimize, time2eval=time2eval) + + analysis = tune.run( + compute_with_config, + config={ + 'x': tune.lograndint(lower=1, upper=1000000), + 'y': tune.randint(lower=1, upper=1000000) + }, + metric='metric2minimize', mode='min', + num_samples=1000000, time_budget_s=60, use_ray=False) + + print(analysis.trials[-1].last_result) + ``` + + Args: + _metric: Optional default anonymous metric for ``tune.report(value)``. + (For compatibility with ray.tune.report) + **kwargs: Any key value pair to be reported. + + Raises: + StopIteration (when not using ray, i.e., _use_ray=False): + A StopIteration exception is raised if the trial has been signaled to stop. + SystemExit (when using ray): + A SystemExit exception is raised if the trial has been signaled to stop by ray. + + A worker/callback thread a trainable spawns during evaluation can call + this too, with no change on the trainable's part: this thread's own + _state.runner is None (a fresh thread never ran tune.run() itself), so + the active run's runner/trial is read from _propagated_context instead + (#996 follow-up, second review point 1). See _RunContext. + + A thread started before tune.run() is called, such as a persistent + queue/callback worker that never calls get_run_context()/ + use_run_context() itself, is attributed automatically when exactly one + tune.run() call is active anywhere in the process right now (#996 + follow-up, eighth review point 1): there is only one run it could mean, + so it is resolved the same way the pre-#996 module-global design + resolved every caller, unconditionally. When zero are active, the + report has nothing to attribute to and is dropped, with a one-time + (per run(), see _warn_context_less_report_once) warning pointing at + get_run_context()/use_run_context(). When two or more are active at + once, which run the report belongs to is genuinely undecidable and this + raises RuntimeError instead of guessing or silently dropping; see + _resolve_unambiguous_active_runner. Either way, this fallback is only + ever a live resolution of "whichever trial this run happens to be + running right now": a worker slow enough to still be draining an item + from an EARLIER trial after its single active run has already moved on + to a later one can still land on the wrong trial, the same risk the + pre-#996 code carried for this exact shape. Callers who need reporting + pinned to the trial that actually produced the item, not whichever one + happens to be current when it is handled, should still capture + get_run_context() where the work is produced and attach it with + use_run_context() where it is processed; that path is unaffected by + this fallback and carries no such risk. + """ + use_ray = _state.use_ray + runner = _state.runner + verbose = _state.verbose + running_trial = None + # True only for the exact precondition the docstring above and + # _warn_context_less_report_once() describe: no thread-local runner and + # no propagated/explicit _RunContext at all (#996 follow-up, seventh + # review). Computed once, before deciding whether to even try ray, + # because it also has to cover the ImportError branch just below: a + # context-less thread's _state.use_ray defaults to True (_TuneState), + # not to whatever the real active run's use_ray actually is, so this + # exact caller shape reaches the ImportError return, not the "no + # trial" one, whenever ray is not installed. Verified directly while + # building this fix: an unchanged pre-started queue worker's report + # never reached the "no trial" branch at all in that environment. + context_less = runner is None + # Guards the finally below: only the fallback branch just below ever + # sets this True, so _state.log_run_id is only ever touched (and only + # ever restored) for exactly the one call that used it. + _restore_log_run_id = False + _prior_log_run_id = None + if runner is None: + ctx = _propagated_context.get() + if ctx is not None: + context_less = False + use_ray = ctx.use_ray + runner = ctx.runner + verbose = ctx.verbose + running_trial = ctx.running_trial + else: + # No thread-local runner, no propagated/explicit context: try + # the single-active-sequential-run fallback (#996 follow-up, + # eighth review point 1) before falling through to the + # use_ray branch below, which would otherwise act on + # _state.use_ray's misleading default (see the comment above) + # instead of the real active run's own backend. Raises + # RuntimeError here, uncaught, when two or more sequential + # runs are active at once, deliberately not folded into the + # try/except ImportError below, which is about ray being + # unavailable, not about ownership being ambiguous. + fallback_runner, fallback_log_run_id, fallback_verbose = _resolve_unambiguous_active_runner() + if fallback_runner is not None: + context_less = False + use_ray = False + runner = fallback_runner + verbose = fallback_verbose + # Scoped to this one report() call only (restored in the + # finally below), the same as use_run_context() scopes it + # to its `with` block: this thread is not "in" the + # resolved run the way a use_run_context() caller + # declares itself to be, only this one dispatch is. + _prior_log_run_id = _state.log_run_id + _state.log_run_id = fallback_log_run_id + _restore_log_run_id = True + try: + if use_ray: + try: + from ray import __version__ as ray_version + + if ray_version.startswith("1."): + from ray import tune + + return tune.report(_metric, **kwargs) + else: # ray>=2 + from ray.air import session + + return session.report(metrics={"metric": _metric, **kwargs}) + except ImportError: + # calling tune.report() outside tune.run(), or (#996 follow-up, + # seventh review) a context-less caller that defaulted here + # instead of to the "no trial" branch below, per the comment + # above. Ray missing means there is nowhere else this call + # could reach either way, so warn under the same condition. + if context_less: + _warn_context_less_report_once() + return + result = kwargs + if _metric is not None: + result[DEFAULT_METRIC] = _metric + # running_trial is the trial a propagated context pinned this report to; + # otherwise (the thread actually driving run()'s own loop) resolve it + # live off the runner, which is always the trial that loop is currently + # stepping. + trial = running_trial if running_trial is not None else getattr(runner, "running_trial", None) + if not trial: + # No thread-local runner and no propagated/explicit context, and + # (see above) no single unambiguous active run to fall back to + # either: this call cannot be attributed to a trial (see the + # docstring above), so it is dropped rather than guessed at. + # context_less can be False here too (a context was found but + # its running_trial had not been resolved yet), which is a + # different, pre-existing edge case this warning is not about, + # so it only fires on the one this review names. + if context_less: + _warn_context_less_report_once() + return None + if trial.is_finished(): + # A late report from a background thread or executor task whose + # captured _RunContext outlived its trial (#996 follow-up, fourth + # review point 2): the trial's final result is already recorded, + # and process_trial_result() would overwrite it with this stale + # value, plus the is_finished() check below would then raise + # StopIteration into a caller that never expected it (unlike the + # trainable's own control-flow loop, which does). Drop it instead. + # This is a fast-path check only, not the admission decision: a + # concurrent report for this same trial can still finish it between + # this line and the lock below, which is what that lock is for. + return None + result["config"] = trial.config + if INCUMBENT_RESULT in result["config"]: + del result["config"][INCUMBENT_RESULT] + for key, value in trial.config.items(): + result["config/" + key] = value + with _admission_lock_for(trial): + # Re-check under the lock (#996 follow-up, fifth review point 2): + # the fast-path check above and this admission are not the same + # instant, and a second, truly concurrent report for this trial + # (a trainable's own worker threads reporting for the trial they + # share, for instance) can legitimately finish it in between. This + # is the only check whose result process_trial_result() actually + # acts on. + if trial.is_finished(): + return None + # Allocated inside this same critical section, immediately before + # the write it orders (#996 follow-up, sixth review point 3): + # _next_training_iteration() used to run before this lock was + # taken, so two truly concurrent reports for this trial could be + # handed iterations in one order (A=5, B=6) and then reach + # process_trial_result() in the OTHER order if B's thread happened + # to acquire the lock first, handing the scheduler/searcher a + # decreasing training_iteration for the trial they track. Locking + # allocation and admission together makes the two always agree: + # whichever report acquires the lock first is both the one that + # gets the lower iteration number and the one process_trial_result() + # sees first. + result["training_iteration"] = _next_training_iteration(trial) + runner.process_trial_result(trial, result) + if verbose > 2: + logger.info(f"result: {result}") + if trial.is_finished(): + raise StopIteration + finally: + if _restore_log_run_id: + _state.log_run_id = _prior_log_run_id + + +def run( + evaluation_function, + config: Optional[dict] = None, + low_cost_partial_config: Optional[dict] = None, + cat_hp_cost: Optional[dict] = None, + metric: Optional[str] = None, + mode: Optional[str] = None, + time_budget_s: Union[int, float] = None, + points_to_evaluate: Optional[List[dict]] = None, + evaluated_rewards: Optional[List] = None, + resource_attr: Optional[str] = None, + min_resource: Optional[float] = None, + max_resource: Optional[float] = None, + reduction_factor: Optional[float] = None, + scheduler=None, + search_alg=None, + verbose: Optional[int] = 2, + local_dir: Optional[str] = None, + num_samples: Optional[int] = 1, + resources_per_trial: Optional[dict] = None, + config_constraints: Optional[List[Tuple[Callable[[dict], float], str, float]]] = None, + metric_constraints: Optional[List[Tuple[str, str, float]]] = None, + max_failure: Optional[int] = 100, + use_ray: Optional[bool] = False, + use_spark: Optional[bool] = False, + use_incumbent_result_in_evaluation: Optional[bool] = None, + log_file_name: Optional[str] = None, + lexico_objectives: Optional[dict] = None, + force_cancel: Optional[bool] = False, + n_concurrent_trials: Optional[int] = 0, + mlflow_exp_name: Optional[str] = None, + automl_info: Optional[Tuple[float]] = None, + extra_tag: Optional[dict] = None, + cost_attr: Optional[str] = "auto", + cost_budget: Optional[float] = None, + **ray_args, +): + """The function-based way of performing HPO. + + Example: + + ```python + import time + from flaml import tune + + def compute_with_config(config): + current_time = time.time() + metric2minimize = (round(config['x'])-95000)**2 + time2eval = time.time() - current_time + tune.report(metric2minimize=metric2minimize, time2eval=time2eval) + # if the evaluation fails unexpectedly and the exception is caught, + # and it doesn't inform the goodness of the config, + # return {} + # if the failure indicates a config is bad, + # report a bad metric value like np.inf or -np.inf + # depending on metric mode being min or max + + analysis = tune.run( + compute_with_config, + config={ + 'x': tune.lograndint(lower=1, upper=1000000), + 'y': tune.randint(lower=1, upper=1000000) + }, + metric='metric2minimize', mode='min', + num_samples=-1, time_budget_s=60, use_ray=False) + + print(analysis.trials[-1].last_result) + ``` + + Args: + evaluation_function: A user-defined evaluation function. + It takes a configuration as input, outputs a evaluation + result (can be a numerical value or a dictionary of string + and numerical value pairs) for the input configuration. + For machine learning tasks, it usually involves training and + scoring a machine learning model, e.g., through validation loss. + config: A dictionary to specify the search space. + low_cost_partial_config: A dictionary from a subset of + controlled dimensions to the initial low-cost values. + e.g., ```{'n_estimators': 4, 'max_leaves': 4}``` + + cat_hp_cost: A dictionary from a subset of categorical dimensions + to the relative cost of each choice. + e.g., ```{'tree_method': [1, 1, 2]}``` + i.e., the relative cost of the + three choices of 'tree_method' is 1, 1 and 2 respectively + metric: A string of the metric name to optimize for. + mode: A string in ['min', 'max'] to specify the objective as + minimization or maximization. + time_budget_s: int or float | The time budget in seconds. + points_to_evaluate: A list of initial hyperparameter + configurations to run first. + evaluated_rewards (list): If you have previously evaluated the + parameters passed in as points_to_evaluate you can avoid + re-running those trials by passing in the reward attributes + as a list so the optimiser can be told the results without + needing to re-compute the trial. Must be the same or shorter length than + points_to_evaluate. + e.g., + + ```python + points_to_evaluate = [ + {"b": .99, "cost_related": {"a": 3}}, + {"b": .99, "cost_related": {"a": 2}}, + ] + evaluated_rewards = [3.0] + ``` + + means that you know the reward for the first config in + points_to_evaluate is 3.0 and want to inform run(). + + resource_attr: A string to specify the resource dimension used by + the scheduler via "scheduler". + min_resource: A float of the minimal resource to use for the resource_attr. + max_resource: A float of the maximal resource to use for the resource_attr. + reduction_factor: A float of the reduction factor used for incremental + pruning. + scheduler: A scheduler for executing the experiment. Can be None, 'flaml', + 'asha' (or 'async_hyperband', 'asynchyperband') or a custom instance of the TrialScheduler class. Default is None: + in this case when resource_attr is provided, the 'flaml' scheduler will be + used, otherwise no scheduler will be used. When set 'flaml', an + authentic scheduler implemented in FLAML will be used. It does not + require users to report intermediate results in evaluation_function. + Find more details about this scheduler in this paper + https://arxiv.org/pdf/1911.04706.pdf). + When set 'asha', the input for arguments "resource_attr", + "min_resource", "max_resource" and "reduction_factor" will be passed + to ASHA's "time_attr", "max_t", "grace_period" and "reduction_factor" + respectively. You can also provide a self-defined scheduler instance + of the TrialScheduler class. When 'asha' or self-defined scheduler is + used, you usually need to report intermediate results in the evaluation + function via 'tune.report()'. + If you would like to do some cleanup opearation when the trial is stopped + by the scheduler, you can catch the `StopIteration` (when not using ray) + or `SystemExit` (when using ray) exception explicitly, + as shown in the following example. + Please find more examples using different types of schedulers + and how to set up the corresponding evaluation functions in + test/tune/test_scheduler.py, and test/tune/example_scheduler.py. + ```python + def easy_objective(config): + width, height = config["width"], config["height"] + for step in range(config["steps"]): + intermediate_score = evaluation_fn(step, width, height) + try: + tune.report(iterations=step, mean_loss=intermediate_score) + except (StopIteration, SystemExit): + # do cleanup operation here + return + ``` + search_alg: An instance/string of the search algorithm + to be used. The same instance can be used for iterative tuning. + e.g., + + ```python + from flaml import BlendSearch + algo = BlendSearch(metric='val_loss', mode='min', + space=search_space, + low_cost_partial_config=low_cost_partial_config) + for i in range(10): + analysis = tune.run(compute_with_config, + search_alg=algo, use_ray=False) + print(analysis.trials[-1].last_result) + ``` + + verbose: 0, 1, 2, or 3. If ray or spark backend is used, their verbosity will be + affected by this argument. 0 = silent, 1 = only status updates, + 2 = status and brief trial results, 3 = status and detailed trial results. + Defaults to 2. + local_dir: A string of the local dir to save ray logs if ray backend is + used; or a local dir to save the tuning log. + num_samples: An integer of the number of configs to try. Defaults to 1. + resources_per_trial: A dictionary of the hardware resources to allocate + per trial, e.g., `{'cpu': 1}`. It is only valid when using ray backend + (by setting 'use_ray = True'). It shall be used when you need to do + [parallel tuning](/docs/Use-Cases/Tune-User-Defined-Function#parallel-tuning). + config_constraints: A list of config constraints to be satisfied. + e.g., ```config_constraints = [(mem_size, '<=', 1024**3)]``` + + mem_size is a function which produces a float number for the bytes + needed for a config. + It is used to skip configs which do not fit in memory. + metric_constraints: A list of metric constraints to be satisfied. + e.g., `['precision', '>=', 0.9]`. The sign can be ">=" or "<=". + max_failure: int | the maximal consecutive number of failures to sample + a trial before the tuning is terminated. + use_ray: A boolean of whether to use ray as the backend. + use_spark: A boolean of whether to use spark as the backend. + log_file_name: A string of the log file name. Default to None. + When set to None: + if local_dir is not given, no log file is created; + if local_dir is given, the log file name will be autogenerated under local_dir. + Only valid when verbose > 0 or use_ray is True. + lexico_objectives: dict, default=None | It specifics information needed to perform multi-objective + optimization with lexicographic preferences. When lexico_objectives is not None, the arguments metric, + mode, will be invalid, and flaml's tune uses CFO + as the `search_alg`, which makes the input (if provided) `search_alg' invalid. + This dictionary shall contain the following fields of key-value pairs: + - "metrics": a list of optimization objectives with the orders reflecting the priorities/preferences of the + objectives. + - "modes" (optional): a list of optimization modes (each mode either "min" or "max") corresponding to the + objectives in the metric list. If not provided, we use "min" as the default mode for all the objectives. + - "targets" (optional): a dictionary to specify the optimization targets on the objectives. The keys are the + metric names (provided in "metric"), and the values are the numerical target values. + - "tolerances" (optional): a dictionary to specify the optimality tolerances on objectives. The keys are the metric names (provided in "metrics"), and the values are the absolute/percentage tolerance in the form of numeric/string. + E.g., + ```python + lexico_objectives = { + "metrics": ["error_rate", "pred_time"], + "modes": ["min", "min"], + "tolerances": {"error_rate": 0.01, "pred_time": 0.0}, + "targets": {"error_rate": 0.0}, + } + ``` + We also support percentage tolerance. + E.g., + ```python + lexico_objectives = { + "metrics": ["error_rate", "pred_time"], + "modes": ["min", "min"], + "tolerances": {"error_rate": "5%", "pred_time": "0%"}, + "targets": {"error_rate": 0.0}, + } + ``` + force_cancel: boolean, default=False | Whether to forcely cancel the PySpark job if overtime. + mlflow_exp_name: str, default=None | The name of the mlflow experiment. This should be specified if + enable mlflow autologging on Spark. Otherwise it will log all the results into the experiment of the + same name as the basename of main entry file. + automl_info: tuple, default=None | The information of the automl run. It should be a tuple of (mlflow_log_latency,). + n_concurrent_trials: int, default=0 | The number of concurrent trials when perform hyperparameter + tuning with Spark. Only valid when use_spark=True and spark is required: + `pip install flaml[spark]`. Please check + [here](https://spark.apache.org/docs/latest/api/python/getting_started/install.html) + for more details about installing Spark. When tune.run() is called from AutoML, it will be + overwritten by the value of `n_concurrent_trials` in AutoML. When <= 0, the concurrent trials + will be set to the number of executors. + extra_tag: dict, default=None | Extra tags to be added to the mlflow runs created by autologging. + cost_attr: None or str to specify the attribute to evaluate the cost of different trials. + Default is "auto", which means that we will automatically choose the cost attribute to use (depending + on the nature of the resource budget). When cost_attr is set to None, cost differences between different trials will be omitted + in our search algorithm. When cost_attr is set to a str different from "auto" and "time_total_s", + this cost_attr must be available in the result dict of the trial. + cost_budget: A float of the cost budget. Only valid when cost_attr is a str different from "auto" and "time_total_s". + **ray_args: keyword arguments to pass to ray.tune.run(). + Only valid when use_ray=True. + """ + global internal_mlflow + old_use_ray = _state.use_ray + old_verbose = _state.verbose + old_runner = _state.runner + old_log_run_id = _state.log_run_id + _run_handler = None + _internal_mlflow = False + mlflow_integration = None + + # Reset once per run() call, not once per process (#996 follow-up, + # eighth review point 2): see the comment above + # _context_less_report_warned's definition for why once-per-run, not + # once-per-report, and why a nested run() on the same thread resetting + # it early is an accepted, safe imprecision. + global _context_less_report_warned + with _context_less_report_warned_lock: + _context_less_report_warned = False + + def _restore_tune_state(): + """Undo every mutation this call made to shared/thread-local state. + + Called from the except clause wrapping the common setup section, + AND from the single outer try/finally that wraps every backend + branch (ray/spark/sequential) below it, so a setup or execution + failure anywhere in this call restores state exactly like a normal + return does, instead of leaking a mutated _state/logger into + whatever run() this thread resumes next. Backend init (the spark + session, `check_spark()`) and the sequential scheduler setup used + to run outside any restoration guard; both are inside the outer + try/finally now (#996 follow-up, second review point 3). + """ + _state.use_ray = old_use_ray + _state.verbose = old_verbose + if not use_ray: + this_runner = _state.runner + _state.runner = old_runner + _state.log_run_id = old_log_run_id + if this_runner is not None: + # Deregister from the single-active-run fallback a + # context-less report() can resolve to (#996 follow-up, + # eighth review point 1): once this call is done restoring + # state, no report anywhere should be able to reach this + # runner anymore, through the fallback or otherwise. Safe + # even if this_runner was never registered (setup failed + # before the SequentialTrialRunner was constructed, or the + # ray/spark branch never touches _active_runners at all). + with _active_runners_lock: + _active_runners.pop(this_runner, None) + if _run_handler is not None: + logger.removeHandler(_run_handler) + _logger_level_exit(_run_handler.level) + if _internal_mlflow: + mlflow_integration.adopt_children() + + try: + if log_file_name: + dir_name = os.path.dirname(log_file_name) + if dir_name: + os.makedirs(dir_name, exist_ok=True) + elif local_dir and verbose > 0: + os.makedirs(local_dir, exist_ok=True) + log_file_name = os.path.join(local_dir, "tune_" + str(datetime.datetime.now()).replace(":", "-") + ".log") + if use_ray and use_spark: + raise ValueError("use_ray and use_spark cannot be both True.") + if not use_ray: + _state.use_ray = False + _state.verbose = verbose + assert not ray_args, "ray_args is only valid when use_ray=True" + _log_run_id = object() + _state.log_run_id = _log_run_id + if verbose > 0: + if log_file_name: + _run_handler = logging.FileHandler(log_file_name) + else: + _run_handler = logging.StreamHandler(stream=sys.stdout) + _run_handler.setFormatter(logger_formatter) + # Filter + addHandler (never `logger.handlers = [...]`) so a + # concurrently active run's handler is never wiped out, and a + # per-run level (not a shared logger.setLevel by itself) so one + # run's verbosity can never silence or tighten another's (#996 + # follow-up: the logger was still a corruptible sixth piece of + # shared state after the five _TuneState fields were fixed). + _run_handler.addFilter(_RunScopedFilter(_log_run_id)) + _run_handler.setLevel(logging.DEBUG if verbose > 2 else logging.INFO) + logger.addHandler(_run_handler) + _logger_level_enter(_run_handler.level) + # verbose == 0 intentionally adds no handler and touches no shared + # state: the old code called logger.setLevel(logging.CRITICAL) here, + # which silenced every OTHER concurrently active run's logging too, + # the same corruption class this whole block now avoids. + + if internal_mlflow and not automl_info and (mlflow.active_run() or is_autolog_enabled()): + mlflow_integration = MLflowIntegration("tune", mlflow_exp_name, extra_tag) + evaluation_function = mlflow_integration.wrap_evaluation_function(evaluation_function) + _internal_mlflow = not automl_info # True if mlflow_integration will be used for logging + else: + _internal_mlflow = False + + from .searcher.blendsearch import CFO, BlendSearch, RandomSearch + + if lexico_objectives is not None: + if "modes" not in lexico_objectives.keys(): + lexico_objectives["modes"] = ["min"] * len(lexico_objectives["metrics"]) + for t_metric, t_mode in zip(lexico_objectives["metrics"], lexico_objectives["modes"]): + if t_metric not in lexico_objectives["tolerances"].keys(): + lexico_objectives["tolerances"][t_metric] = 0 + if t_metric not in lexico_objectives["targets"].keys(): + lexico_objectives["targets"][t_metric] = -float("inf") if t_mode == "min" else float("inf") + if search_alg is None or isinstance(search_alg, str): + if isinstance(search_alg, str): + assert search_alg in [ + "BlendSearch", + "CFO", + "CFOCat", + "RandomSearch", + ], f"search_alg={search_alg} is not recognized. 'BlendSearch', 'CFO', 'CFOcat' and 'RandomSearch' are supported." + + flaml_scheduler_resource_attr = ( + flaml_scheduler_min_resource + ) = flaml_scheduler_max_resource = flaml_scheduler_reduction_factor = None + if scheduler in (None, "flaml"): + # when scheduler is set 'flaml' or None, we will use a scheduler that is + # authentic to the search algorithms in flaml. After setting up + # the search algorithm accordingly, we need to set scheduler to + # None in case it is later used in the trial runner. + flaml_scheduler_resource_attr = resource_attr + flaml_scheduler_min_resource = min_resource + flaml_scheduler_max_resource = max_resource + flaml_scheduler_reduction_factor = reduction_factor + scheduler = None + if lexico_objectives: + # TODO: Modify after supporting BlendSearch in lexicographic optimization + SearchAlgorithm = CFO + logger.info( + f"Using search algorithm {SearchAlgorithm.__name__} for lexicographic optimization. Note that when providing other search algorithms, we use CFO instead temporarily." + ) + metric = lexico_objectives["metrics"][0] or DEFAULT_METRIC + else: + if not search_alg or search_alg == "BlendSearch": + try: + import optuna as _ + + SearchAlgorithm = BlendSearch + logger.info(f"Using search algorithm {SearchAlgorithm.__name__}.") + except ImportError: + if search_alg == "BlendSearch": + raise ValueError("To use BlendSearch, run: pip install flaml[blendsearch]") + else: + SearchAlgorithm = CFO + logger.warning( + "Using CFO for search. To use BlendSearch, run: pip install flaml[blendsearch]" + ) + else: + SearchAlgorithm = locals()[search_alg] + logger.info(f"Using search algorithm {SearchAlgorithm.__name__}.") + metric = metric or DEFAULT_METRIC + search_alg = SearchAlgorithm( + metric=metric, + mode=mode, + space=config, + points_to_evaluate=points_to_evaluate, + evaluated_rewards=evaluated_rewards, + low_cost_partial_config=low_cost_partial_config, + cat_hp_cost=cat_hp_cost, + time_budget_s=time_budget_s, + num_samples=num_samples, + resource_attr=flaml_scheduler_resource_attr, + min_resource=flaml_scheduler_min_resource, + max_resource=flaml_scheduler_max_resource, + reduction_factor=flaml_scheduler_reduction_factor, + config_constraints=config_constraints, + metric_constraints=metric_constraints, + use_incumbent_result_in_evaluation=use_incumbent_result_in_evaluation, + lexico_objectives=lexico_objectives, + cost_attr=cost_attr, + cost_budget=cost_budget, + ) + else: + if metric is None or mode is None: + if lexico_objectives: + metric = lexico_objectives["metrics"][0] or metric or search_alg.metric or DEFAULT_METRIC + mode = lexico_objectives["modes"][0] or mode or search_alg.mode + else: + metric = metric or search_alg.metric or DEFAULT_METRIC + mode = mode or search_alg.mode + if ray_available and use_ray: + if ray_version.startswith("1."): + from ray.tune.suggest import ConcurrencyLimiter + else: + from ray.tune.search import ConcurrencyLimiter + else: + from flaml.tune.searcher.suggestion import ConcurrencyLimiter + if ( + search_alg.__class__.__name__ + in [ + "BlendSearch", + "CFO", + "CFOCat", + ] + and use_incumbent_result_in_evaluation is not None + ): + search_alg.use_incumbent_result_in_evaluation = use_incumbent_result_in_evaluation + searcher = search_alg.searcher if isinstance(search_alg, ConcurrencyLimiter) else search_alg + if lexico_objectives: + # TODO: Modify after supporting BlendSearch in lexicographic optimization + assert search_alg.__class__.__name__ in [ + "CFO", + ], "If lexico_objectives is not None, the search_alg must be CFO for now." + search_alg.lexico_objective = lexico_objectives + + if isinstance(searcher, BlendSearch): + setting = {} + if time_budget_s: + setting["time_budget_s"] = time_budget_s + if num_samples > 0: + setting["num_samples"] = num_samples + searcher.set_search_properties(metric, mode, config, **setting) + else: + searcher.set_search_properties(metric, mode, config) + if scheduler in ("asha", "asynchyperband", "async_hyperband"): + params = {} + # scheduler resource_dimension=resource_attr + if resource_attr: + params["time_attr"] = resource_attr + if max_resource: + params["max_t"] = max_resource + if min_resource: + params["grace_period"] = min_resource + if reduction_factor: + params["reduction_factor"] = reduction_factor + if ray_available: + from ray.tune.schedulers import ASHAScheduler + + scheduler = ASHAScheduler(**params) + except Exception: + _restore_tune_state() + raise + + # One outer try/finally for every backend branch below (ray, spark, + # sequential): Spark/backend initialization (check_spark(), the + # SparkSession, register_spark()) and the sequential path's scheduler + # setup used to run before any of these branches' own try/finally + # started, so a failure there raised straight out of run() without + # calling _restore_tune_state() at all, leaking the earlier setup + # section's _state/logger mutations into whatever this thread runs + # next, nested tune.run() call or caught-and-retried one alike (#996 + # follow-up, second review point 3). Wrapping from here means every + # setup step and every execution branch shares the one guard. + try: + if use_ray: + try: + from ray import tune + except ImportError: + raise ImportError("Failed to import ray tune. " "Please install ray[tune] or set use_ray=False") + _state.use_ray = True + analysis = tune.run( + evaluation_function, + metric=metric, + mode=mode, + search_alg=search_alg, + scheduler=scheduler, + time_budget_s=time_budget_s, + verbose=verbose, + local_dir=local_dir, + num_samples=num_samples, + resources_per_trial=resources_per_trial, + **ray_args, + ) + if log_file_name: + with open(log_file_name, "w") as f: + for trial in analysis.trials: + f.write(f"result: {trial.last_result}\n") + return analysis + + if use_spark: + # parallel run with spark + spark_available, spark_error_msg = check_spark() + if not spark_available: + raise spark_error_msg + try: + from joblib import Parallel, delayed, parallel_backend + from joblibspark import register_spark + from pyspark.sql import SparkSession + except ImportError as e: + raise ImportError(f"{e}. Try pip install flaml[spark] or set use_spark=False.") + from flaml.tune.searcher.suggestion import ConcurrencyLimiter + + from .trial_runner import SparkTrialRunner + + register_spark() + spark = SparkSession.builder.getOrCreate() + sc = spark._jsc.sc() + num_executors = len([executor.host() for executor in sc.statusTracker().getExecutorInfos()]) - 1 + """ + By default, the number of executors is the number of VMs in the cluster. And we can + launch one trial per executor. However, sometimes we can launch more trials than + the number of executors (e.g., local mode). In this case, we can set the environment + variable `FLAML_MAX_CONCURRENT` to override the detected `num_executors`. + + `max_concurrent` is the maximum number of concurrent trials defined by `search_alg`, + `FLAML_MAX_CONCURRENT` will also be used to override `max_concurrent` if `search_alg` + is not an instance of `ConcurrencyLimiter`. + + The final number of concurrent trials is the minimum of `max_concurrent` and + `num_executors` if `n_concurrent_trials<=0` (default, automl cases), otherwise the + minimum of `max_concurrent` and `n_concurrent_trials` (tuning cases). + """ + time_start = time.time() + try: + FLAML_MAX_CONCURRENT = int(os.getenv("FLAML_MAX_CONCURRENT", 0)) + except ValueError: + FLAML_MAX_CONCURRENT = 0 + num_executors = max(num_executors, FLAML_MAX_CONCURRENT, 1) + max_spark_parallelism = max(spark.sparkContext.defaultParallelism, FLAML_MAX_CONCURRENT) + if scheduler: + scheduler.set_search_properties(metric=metric, mode=mode) + if isinstance(search_alg, ConcurrencyLimiter): + max_concurrent = max(1, search_alg.max_concurrent) + else: + max_concurrent = max(1, max_spark_parallelism) + passed_in_n_concurrent_trials = max(n_concurrent_trials, max_concurrent) + n_concurrent_trials = min( + n_concurrent_trials if n_concurrent_trials > 0 else num_executors, + max_concurrent, + ) + if n_concurrent_trials < passed_in_n_concurrent_trials: + logger.warning( + f"The actual concurrent trials is {n_concurrent_trials}. You can set the environment " + f"variable `FLAML_MAX_CONCURRENT` to '{passed_in_n_concurrent_trials}' to override the detected num of executors." + ) + with parallel_backend("spark"): + with Parallel(n_jobs=n_concurrent_trials, verbose=max(0, (verbose - 1) * 50)) as parallel: + _state.runner = SparkTrialRunner( + search_alg=search_alg, + scheduler=scheduler, + metric=metric, + mode=mode, + ) + num_trials = 0 + if time_budget_s is None: + time_budget_s = np.inf + num_failures = 0 + upperbound_num_failures = (len(evaluated_rewards) if evaluated_rewards else 0) + max_failure + logger.debug(f"automl_info: {automl_info}") + while ( + time.time() - time_start < time_budget_s + and (num_samples < 0 or num_trials < num_samples) + and num_failures < upperbound_num_failures + ): + if automl_info and automl_info[1] == "all" and automl_info[0] > 0 and time_budget_s < np.inf: + time_budget_s -= automl_info[0] * n_concurrent_trials + logger.debug(f"Remaining time budget with mlflow log latency: {time_budget_s} seconds.") + while len(_state.runner.running_trials) < n_concurrent_trials: + # suggest trials for spark + trial_next = _state.runner.step() + if trial_next: + num_trials += 1 + else: + num_failures += 1 # break with upperbound_num_failures consecutive failures + logger.debug(f"consecutive failures is {num_failures}") + if num_failures >= upperbound_num_failures: + break + trials_to_run = _state.runner.running_trials + if not trials_to_run: + logger.warning(f"fail to sample a trial for {max_failure} times in a row, stopping.") + break + logger.info( + f"Number of trials: {num_trials}/{num_samples}, {len(_state.runner.running_trials)} RUNNING," + f" {len(_state.runner._trials) - len(_state.runner.running_trials)} TERMINATED" + ) + logger.debug( + f"Configs of Trials to run: {[trial_to_run.config for trial_to_run in trials_to_run]}" + ) + results = None + with PySparkOvertimeMonitor(time_start, time_budget_s, force_cancel, parallel=parallel): + try: + results = parallel( + delayed(evaluation_function)(trial_to_run.config) for trial_to_run in trials_to_run + ) + except RuntimeError as e: + logger.warning(f"RuntimeError: {e}") + results = None + logger.info( + "Encountered RuntimeError. Waiting 10 seconds for Spark cluster to recover before retrying." + ) + time.sleep(10) + # results = [evaluation_function(trial_to_run.config) for trial_to_run in trials_to_run] + while results: + result = results.pop(0) + trial_to_run = trials_to_run[0] + _state.runner.running_trial = trial_to_run + if result is not None: + if _internal_mlflow: + mlflow_integration.record_trial(result, trial_to_run, metric) + + if isinstance(result, dict): + if result: + logger.info(f"Brief result: {result}") + report(**result) + else: + # When the result returned is an empty dict, set the trial status to error + trial_to_run.set_status(Trial.ERROR) + else: + logger.info("Brief result: {metric: result}") + report(_metric=result) + _state.runner.stop_trial(trial_to_run) + num_failures = 0 + analysis = ExperimentAnalysis( + _state.runner.get_trials(), + metric=metric, + mode=mode, + lexico_objectives=lexico_objectives, + ) + analysis.search_space = config + + if _internal_mlflow: + mlflow_integration.log_tune(analysis, metric) + # try: + # _best_config = analysis.best_config + # except Exception: + # _best_config = None + # if _best_config: + # parallel( + # delayed(mlflow_integration.retrain)(evaluation_function, analysis.best_config) + # for dummy in [0] + # ) + + return analysis + + # simple sequential run without using tune.run() from ray + time_start = time.time() + _state.use_ray = False + if scheduler: + scheduler.set_search_properties(metric=metric, mode=mode) + from .trial_runner import SequentialTrialRunner + + _state.runner = SequentialTrialRunner( + search_alg=search_alg, + scheduler=scheduler, + metric=metric, + mode=mode, + ) + # Registers this run as the (so far) sole candidate a context-less + # report() can fall back to (#996 follow-up, eighth review point + # 1); see _resolve_unambiguous_active_runner. Deregistered in + # _restore_tune_state() above, which every exit path from here + # (normal return, break, or an exception caught by the outer + # try/finally) runs. + with _active_runners_lock: + _active_runners[_state.runner] = (_state.log_run_id, verbose) + num_trials = 0 + if time_budget_s is None: + time_budget_s = np.inf + num_failures = 0 + upperbound_num_failures = (len(evaluated_rewards) if evaluated_rewards else 0) + max_failure + while ( + time.time() - time_start < time_budget_s + and (num_samples < 0 or num_trials < num_samples) + and num_failures < upperbound_num_failures + ): + trial_to_run = _state.runner.step() + if trial_to_run: + num_trials += 1 + if verbose: + logger.info(f"trial {num_trials} config: {trial_to_run.config}") + result = None + # Pin this evaluation call's runner/trial/log_run_id in + # _propagated_context so a worker thread evaluation_function + # spawns on its own (or a task it submits to a + # ThreadPoolExecutor) can call tune.report() and land on the + # right trial, with its log records landing in the right + # run log, with no trainable-side change (#996 follow-up, + # second review point 1; log_run_id: third review point 4). + # _install_thread_context_propagation() and + # _install_executor_context_propagation() are what make a + # plain threading.Thread or executor task started inside + # this call see it. + _prop_token = _propagated_context.set( + _RunContext(_state.use_ray, _state.runner, _state.verbose, trial_to_run, _state.log_run_id) + ) + try: + with PySparkOvertimeMonitor(time_start, time_budget_s, force_cancel): + result = evaluation_function(trial_to_run.config) + finally: + _propagated_context.reset(_prop_token) + logger.debug(f"result in tune: {trial_to_run}, {result}") + if result is not None: + if _internal_mlflow: + mlflow_integration.record_trial(result, trial_to_run, metric) + + if isinstance(result, dict): + if result: + report(**result) + else: + # When the result returned is an empty dict, set the trial status to error + trial_to_run.set_status(Trial.ERROR) + else: + report(_metric=result) + _state.runner.stop_trial(trial_to_run) + num_failures = 0 + if trial_to_run.last_result is None: + # application stops tuning by returning None + # TODO document this feature when it is finalized + break + else: + # break with upperbound_num_failures consecutive failures + num_failures += 1 + if num_failures == upperbound_num_failures: + logger.warning(f"fail to sample a trial for {max_failure} times in a row, stopping.") + analysis = ExperimentAnalysis( + _state.runner.get_trials(), + metric=metric, + mode=mode, + lexico_objectives=lexico_objectives, + ) + analysis.search_space = config + if _internal_mlflow: + mlflow_integration.log_tune(analysis, metric) + if analysis.best_run_id is not None: + logger.info(f"Best MLflow run name: {analysis.best_run_name}") + logger.info(f"Best MLflow run id: {analysis.best_run_id}") + # try: + # _best_config = analysis.best_config + # except Exception: + # _best_config = None + # if _best_config: + # mlflow_integration.retrain(evaluation_function, analysis.best_config) + + return analysis + finally: + # recover the global/shared state in case of nested run, or a + # failure anywhere in the block above (#996 follow-up point 3) + _restore_tune_state() + + +class Tuner: + """Tuner is the class-based way of launching hyperparameter tuning jobs compatible with Ray Tune 2. + + Args: + trainable: A user-defined evaluation function. + It takes a configuration as input, outputs a evaluation + result (can be a numerical value or a dictionary of string + and numerical value pairs) for the input configuration. + For machine learning tasks, it usually involves training and + scoring a machine learning model, e.g., through validation loss. + param_space: Search space of the tuning job. + One thing to note is that both preprocessor and dataset can be tuned here. + tune_config: Tuning algorithm specific configs. + Refer to ray.tune.tune_config.TuneConfig for more info. + run_config: Runtime configuration that is specific to individual trials. + If passed, this will overwrite the run config passed to the Trainer, + if applicable. Refer to ray.air.config.RunConfig for more info. + + Usage pattern: + + .. code-block:: python + + from sklearn.datasets import load_breast_cancer + + from ray import tune + from ray.data import from_pandas + from ray.air.config import RunConfig, ScalingConfig + from ray.train.xgboost import XGBoostTrainer + from ray.tune.tuner import Tuner + + def get_dataset(): + data_raw = load_breast_cancer(as_frame=True) + dataset_df = data_raw["data"] + dataset_df["target"] = data_raw["target"] + dataset = from_pandas(dataset_df) + return dataset + + trainer = XGBoostTrainer( + label_column="target", + params={}, + datasets={"train": get_dataset()}, + ) + + param_space = { + "scaling_config": ScalingConfig( + num_workers=tune.grid_search([2, 4]), + resources_per_worker={ + "CPU": tune.grid_search([1, 2]), + }, + ), + # You can even grid search various datasets in Tune. + # "datasets": { + # "train": tune.grid_search( + # [ds1, ds2] + # ), + # }, + "params": { + "objective": "binary:logistic", + "tree_method": "approx", + "eval_metric": ["logloss", "error"], + "eta": tune.loguniform(1e-4, 1e-1), + "subsample": tune.uniform(0.5, 1.0), + "max_depth": tune.randint(1, 9), + }, + } + tuner = Tuner(trainable=trainer, param_space=param_space, + run_config=RunConfig(name="my_tune_run")) + analysis = tuner.fit() + + To retry a failed tune run, you can then do + + .. code-block:: python + + tuner = Tuner.restore(experiment_checkpoint_dir) + tuner.fit() + + ``experiment_checkpoint_dir`` can be easily located near the end of the + console output of your first failed run. + """ diff --git a/test/tune/test_concurrent_run.py b/test/tune/test_concurrent_run.py new file mode 100644 index 0000000000..db8a13f63e --- /dev/null +++ b/test/tune/test_concurrent_run.py @@ -0,0 +1,1292 @@ +"""Regression test for #996: two threads calling flaml.tune.run() concurrently +must not corrupt each other's trial state. + +flaml.tune.run()/tune.report() used to coordinate through module-level globals +(_runner/_running_trial/_training_iteration/_use_ray/_verbose in flaml/tune/tune.py). +Those globals are shared by every thread, so whichever thread's run() call +executes _runner = SequentialTrialRunner(...) most recently "wins" the global +for every thread's subsequent report()/stop_trial() calls, including resetting +it to None (or another thread's runner) out from under a still-running call. + +This forces the exact interleaving that corrupts state, deterministically (no +sleep-based timing): thread A steps its own trial and pauses inside its +evaluation function; thread B starts while A is paused, creates its own runner +(overwriting the shared state A is still relying on), steps its own trial, and +also pauses. A is released and runs to completion, all the way through +report(), stop_trial(), building its analysis, and its own finally-restore, +entirely while B is still paused. B is then released: on main, B's report() +silently drops its result +(the runner it reads has already been reset to whatever A's finally restored), +and B's stop_trial() call raises AttributeError on that stale/None runner. +""" + +import concurrent.futures +import logging +import queue +import threading +from unittest import mock + +import pytest + +import flaml.tune.tune as tune_module +from flaml import tune +from flaml.tune.logger import logger +from flaml.tune.trial import Trial + + +def test_concurrent_tune_run_does_not_corrupt_state(): + a_paused = threading.Event() + b_paused = threading.Event() + release_a = threading.Event() + release_b = threading.Event() + a_done = threading.Event() + results = {} + + def eval_a(config): + a_paused.set() + assert release_a.wait(timeout=5), "thread A was never released" + return {"metric": 1.0} + + def eval_b(config): + b_paused.set() + assert release_b.wait(timeout=5), "thread B was never released" + return {"metric": 2.0} + + def run_a(): + try: + analysis = tune.run( + eval_a, + config={"x": tune.uniform(0, 1)}, + metric="metric", + mode="min", + num_samples=1, + verbose=0, + ) + results["A"] = ("ok", [t.last_result for t in analysis.trials]) + except Exception as e: # noqa: BLE001 (captured for the test assertion below) + results["A"] = ("error", f"{type(e).__name__}: {e}") + finally: + a_done.set() + + def run_b(): + try: + analysis = tune.run( + eval_b, + config={"x": tune.uniform(0, 1)}, + metric="metric", + mode="min", + num_samples=1, + verbose=0, + ) + results["B"] = ("ok", [t.last_result for t in analysis.trials]) + except Exception as e: # noqa: BLE001 (captured for the test assertion below) + results["B"] = ("error", f"{type(e).__name__}: {e}") + + thread_a = threading.Thread(target=run_a) + thread_b = threading.Thread(target=run_b) + + thread_a.start() + assert a_paused.wait(timeout=5), "thread A never reached its evaluation function" + + thread_b.start() + assert b_paused.wait(timeout=5), "thread B never reached its evaluation function" + + # Let A run to full completion, including its own finally-restore, while + # B is still paused inside its evaluation function. + release_a.set() + assert a_done.wait(timeout=5), "thread A never finished" + thread_a.join(timeout=5) + + release_b.set() + thread_b.join(timeout=5) + + def metric_of(key): + status, trials = results.get(key, ("missing", None)) + if status != "ok" or trials is None or len(trials) != 1: + return None + return trials[0].get("metric") + + assert results.get("A", ("missing",))[0] == "ok", f"thread A did not complete cleanly: {results.get('A')}" + assert results.get("B", ("missing",))[0] == "ok", f"thread B did not complete cleanly: {results.get('B')}" + assert metric_of("A") == 1.0, f"thread A's own trial did not receive thread A's own metric: {results}" + assert metric_of("B") == 2.0, f"thread B's own trial did not receive thread B's own metric: {results}" + + +def test_concurrent_tune_run_logging_does_not_cross_contaminate(tmp_path): + """Follow-up to #996: the five _TuneState fields are per-thread now, but + tune.run() was still swapping the shared flaml.tune.logger logger's + `.handlers`/level wholesale (`logger.handlers = []`, then re-add), which + is the same corruption class on a sixth piece of shared state. Two + concurrent tune.run() calls with different log_file_name used to lose + each other's FileHandler mid-run and restore whichever handler list + happened to be current when each one's finally block ran. + + Forces the same deterministic pause/release overlap as + test_concurrent_tune_run_does_not_corrupt_state above: both threads have + their own FileHandler simultaneously attached to the shared logger at + the same time, not just simultaneously "in tune.run()". + """ + a_paused = threading.Event() + b_paused = threading.Event() + release_a = threading.Event() + release_b = threading.Event() + + log_a = str(tmp_path / "a.log") + log_b = str(tmp_path / "b.log") + + handlers_before = list(logger.handlers) + level_before = logger.getEffectiveLevel() + + def eval_a(config): + logger.info("MARKER_FROM_A_PRE") + a_paused.set() + assert release_a.wait(timeout=5), "thread A was never released" + logger.info("MARKER_FROM_A_POST") + return {"metric": 1.0} + + def eval_b(config): + logger.info("MARKER_FROM_B_PRE") + b_paused.set() + assert release_b.wait(timeout=5), "thread B was never released" + logger.info("MARKER_FROM_B_POST") + return {"metric": 2.0} + + def run_a(): + tune.run( + eval_a, + config={"x": tune.uniform(0, 1)}, + metric="metric", + mode="min", + num_samples=1, + verbose=2, + log_file_name=log_a, + ) + + def run_b(): + tune.run( + eval_b, + config={"x": tune.uniform(0, 1)}, + metric="metric", + mode="min", + num_samples=1, + verbose=2, + log_file_name=log_b, + ) + + thread_a = threading.Thread(target=run_a) + thread_b = threading.Thread(target=run_b) + + thread_a.start() + assert a_paused.wait(timeout=5), "thread A never reached its evaluation function" + + thread_b.start() + assert b_paused.wait(timeout=5), "thread B never reached its evaluation function" + + # Both threads' own FileHandlers are attached to the shared logger right + # now. Release A first and let it run to full completion, including its + # own finally-restore, while B is still paused: a proper LIFO nesting + # (innermost starts and finishes first) restores correctly even with the + # old wholesale logger.handlers swap, so this crossing order is the one + # that actually exercises the corruption (same shape as + # test_concurrent_tune_run_does_not_corrupt_state above). + release_a.set() + thread_a.join(timeout=5) + release_b.set() + thread_b.join(timeout=5) + + text_a = open(log_a).read() + text_b = open(log_b).read() + + for marker in ("MARKER_FROM_A_PRE", "MARKER_FROM_A_POST"): + assert marker in text_a, f"thread A's own log file is missing {marker}: {text_a!r}" + assert marker not in text_b, f"{marker} leaked into thread B's log file: {text_b!r}" + for marker in ("MARKER_FROM_B_PRE", "MARKER_FROM_B_POST"): + assert marker in text_b, f"thread B's own log file is missing {marker}: {text_b!r}" + assert marker not in text_a, f"{marker} leaked into thread A's log file: {text_a!r}" + + assert logger.handlers == handlers_before, f"logger.handlers was not fully restored: {logger.handlers}" + assert ( + logger.getEffectiveLevel() == level_before + ), f"logger level was not restored: {logger.getEffectiveLevel()} != {level_before}" + + +def test_tune_run_setup_failure_restores_state(): + """Follow-up to #996: _state.use_ray/_state.verbose and the logger + handler/level are mutated by the `if not use_ray:` block near the top of + run(), before the try/finally that used to be the only thing restoring + them. A failure later in setup (searcher/scheduler construction) used to + skip straight past that try/finally and leak the mutated state into + whatever tune.run() this thread calls next, nested or not. + + `search_alg="not-a-real-search-alg"` fails the validity assert inside + run()'s searcher setup, after the logger/state mutation has already + happened, which is exactly the gap being tested. + """ + handlers_before = list(logger.handlers) + level_before = logger.getEffectiveLevel() + use_ray_before = tune.tune._state.use_ray + verbose_before = tune.tune._state.verbose + + with pytest.raises(AssertionError): + tune.run( + lambda config: {"metric": 1.0}, + config={"x": tune.uniform(0, 1)}, + metric="metric", + mode="min", + num_samples=1, + verbose=2, + search_alg="not-a-real-search-alg", + ) + + assert logger.handlers == handlers_before, f"logger.handlers leaked past the setup failure: {logger.handlers}" + assert logger.getEffectiveLevel() == level_before, "logger level leaked past the setup failure" + assert tune.tune._state.use_ray == use_ray_before, "_state.use_ray leaked past the setup failure" + assert tune.tune._state.verbose == verbose_before, "_state.verbose leaked past the setup failure" + + # And state genuinely was not left corrupted: an ordinary run right after + # still works, rather than inheriting whatever the failed setup left behind. + analysis = tune.run( + lambda config: {"metric": 1.0}, + config={"x": tune.uniform(0, 1)}, + metric="metric", + mode="min", + num_samples=1, + verbose=0, + ) + assert len(analysis.trials) == 1 + assert analysis.trials[0].last_result.get("metric") == 1.0 + + +def test_tune_report_from_trainable_spawned_thread(): + """Follow-up to #996, reviewer point 1 (second review): report() reads + _state.runner, which is thread-local (per-thread by design, so + concurrent tune.run() calls do not see each other's runner). A + worker/callback thread that a trainable spawns on its own therefore + starts from fresh _TuneState defaults, with no runner attached. + + An earlier version of this fix required the trainable to call + get_run_context()/use_run_context() itself, and the reviewer asked for + that to work automatically instead, for existing trainables that never + call either. First half now checks exactly that: a plain + threading.Thread with no FLAML-specific code in it still reports + correctly, because run() sets _propagated_context around the + evaluation_function() call and a patched threading.Thread.start() + carries that ambient context into any thread spawned during it (see + _install_thread_context_propagation() in tune.py). Second half checks + the explicit get_run_context()/use_run_context() API still works too, + for callers who want it (a persistent thread-pool executor, for + instance, where the automatic patch can't reach individual submissions). + """ + + def eval_automatic_propagation(config): + def worker(): + tune.report(metric=42.0) + + t = threading.Thread(target=worker) + t.start() + t.join(timeout=5) + return None # the worker thread's report() is the only result + + analysis = tune.run( + eval_automatic_propagation, + config={"x": tune.uniform(0, 1)}, + metric="metric", + mode="min", + num_samples=1, + verbose=0, + ) + assert analysis.trials[0].last_result is not None, ( + "report() from a plain worker thread with no explicit context call was dropped; " + "expected automatic propagation via _propagated_context" + ) + assert analysis.trials[0].last_result.get("metric") == 42.0 + + def eval_with_explicit_propagation(config): + ctx = tune.get_run_context() + + def worker(): + with tune.use_run_context(ctx): + tune.report(metric=7.0) + + t = threading.Thread(target=worker) + t.start() + t.join(timeout=5) + return None + + analysis2 = tune.run( + eval_with_explicit_propagation, + config={"x": tune.uniform(0, 1)}, + metric="metric", + mode="min", + num_samples=1, + verbose=0, + ) + assert analysis2.trials[0].last_result is not None, "propagated report() from the worker thread was dropped" + assert analysis2.trials[0].last_result.get("metric") == 7.0 + + +def test_training_iteration_shared_across_worker_threads(): + """Follow-up to #996, reviewer point 2 (second review): training_iteration + used to be a plain int copied onto _RunContext by get_run_context() and + discarded when the receiving thread's use_run_context() block exited. + The driving thread's own copy was never updated by a worker thread's + report(), so each new propagated handoff restarted counting from + whatever the driving thread's stale copy held (0 here, since the + driving thread itself never reports directly) instead of continuing + the trial's real count. Every one of the three handoffs below would + report training_iteration=1 on unfixed code, not 0, 1, 2. + + Three SEPARATE worker threads report for the SAME trial, one at a + time (joined before the next starts, so this is deterministic, not a + race): each is a brand-new threading.Thread with its own fresh + _TuneState, so this exercises whether the iteration counter is + actually shared through the trial, not just accidentally continuous + because it stayed on one thread. + """ + + def eval_multi_handoff(config): + ctx = tune.get_run_context() + for _ in range(3): + + def worker(): + with tune.use_run_context(ctx): + tune.report(metric=1.0) + + t = threading.Thread(target=worker) + t.start() + t.join(timeout=5) + return None + + analysis = tune.run( + eval_multi_handoff, + config={"x": tune.uniform(0, 1)}, + metric="metric", + mode="min", + num_samples=1, + verbose=0, + ) + last_result = analysis.trials[0].last_result + assert last_result is not None + assert last_result.get("training_iteration") == 2, ( + "expected the third propagated report to record training_iteration=2 (0-indexed, " + f"monotonically increasing across the three handoffs); got {last_result}" + ) + + +def test_tune_run_spark_setup_failure_restores_state(): + """Follow-up to #996, reviewer point 3 (second review): Spark backend + initialization (check_spark(), constructing the SparkSession) used to + run before the try/finally that calls _restore_tune_state(), so a + failure there (here: forced via check_spark(), the same failure a user + hits from a bad environment) raised straight out of run() without + restoring the _state/logger mutations the earlier common-setup section + had already made, leaking them into whatever this thread does next. + + Patches check_spark() itself rather than relying on PySpark being + absent: CI installs real pyspark on some legs, where check_spark() + would otherwise succeed and this test would never exercise the + failure path at all. + """ + handlers_before = list(logger.handlers) + level_before = logger.getEffectiveLevel() + use_ray_before = tune.tune._state.use_ray + verbose_before = tune.tune._state.verbose + + with mock.patch( + "flaml.tune.tune.check_spark", + return_value=(False, ImportError("simulated: pyspark unavailable")), + ): + with pytest.raises(ImportError): + tune.run( + lambda config: {"metric": 1.0}, + config={"x": tune.uniform(0, 1)}, + metric="metric", + mode="min", + num_samples=1, + verbose=2, + use_spark=True, + ) + + assert logger.handlers == handlers_before, f"logger.handlers leaked past the spark setup failure: {logger.handlers}" + assert logger.getEffectiveLevel() == level_before, "logger level leaked past the spark setup failure" + assert tune.tune._state.use_ray == use_ray_before, "_state.use_ray leaked past the spark setup failure" + assert tune.tune._state.verbose == verbose_before, "_state.verbose leaked past the spark setup failure" + + # State genuinely was not left corrupted: an ordinary run right after + # still works, rather than inheriting whatever the failed setup left. + analysis = tune.run( + lambda config: {"metric": 1.0}, + config={"x": tune.uniform(0, 1)}, + metric="metric", + mode="min", + num_samples=1, + verbose=0, + ) + assert len(analysis.trials) == 1 + assert analysis.trials[0].last_result.get("metric") == 1.0 + + +def test_tune_report_from_prewarmed_threadpool_executor_worker(): + """Follow-up to #996, third review point 1: a ThreadPoolExecutor's + worker threads call Thread.start() once, when the pool spins them up, + not once per submitted task. _install_thread_context_propagation() + captures the ambient _RunContext at start() time, so a worker warmed up + BEFORE any tune.run() call exists would previously capture nothing, and + every later task submitted to that same (reused) worker would silently + lose report() the same way an un-patched plain Thread used to. + + The pool is created and its worker warmed up (submit + result()) before + tune.run() is ever called, so this cannot pass by accident just because + the worker happened to start during a run. + """ + executor = concurrent.futures.ThreadPoolExecutor(max_workers=1) + executor.submit(lambda: None).result() # warm up the one worker thread + + def eval_via_prewarmed_pool(config): + future = executor.submit(tune.report, metric=99.0) + future.result(timeout=5) + return None # the pool worker's report() is the only result + + try: + analysis = tune.run( + eval_via_prewarmed_pool, + config={"x": tune.uniform(0, 1)}, + metric="metric", + mode="min", + num_samples=1, + verbose=0, + ) + finally: + executor.shutdown(wait=True) + + assert analysis.trials[0].last_result is not None, ( + "report() submitted to an already-warmed-up ThreadPoolExecutor worker was dropped; " + "expected _install_executor_context_propagation() to attach the context per task" + ) + assert analysis.trials[0].last_result.get("metric") == 99.0 + + +def test_tune_report_from_prewarmed_executor_reused_across_trials(): + """Same mechanism as above, exercised across TWO trials sharing the SAME + pre-warmed worker thread (the "cross-trial executor reuse" half of + third review point 1): each trial's task must land on ITS OWN trial, + not get pinned to whichever trial warmed the worker up first. + """ + executor = concurrent.futures.ThreadPoolExecutor(max_workers=1) + executor.submit(lambda: None).result() + + def eval_via_prewarmed_pool(config): + future = executor.submit(tune.report, metric=config["x"]) + future.result(timeout=5) + return None + + try: + analysis = tune.run( + eval_via_prewarmed_pool, + config={"x": tune.uniform(0, 1)}, + points_to_evaluate=[{"x": 11.0}, {"x": 22.0}], + metric="metric", + mode="min", + num_samples=2, + verbose=0, + ) + finally: + executor.shutdown(wait=True) + + reported = sorted(t.last_result.get("metric") for t in analysis.trials if t.last_result is not None) + assert reported == [ + 11.0, + 22.0, + ], f"expected each of the two trials to land its own report() through the reused pool worker, got {reported}" + + +def test_tune_report_from_thread_subclass_overriding_run(): + """Follow-up to #996, third review point 2: the first version of this + patch replaced threading.Thread.run directly, which a Thread SUBCLASS + overriding its own run() (idiomatic, common) shadows, so the patch's + context-attaching code never executed for it. The fix wraps + Thread._bootstrap_inner instead, which CPython calls internally and + which invokes self.run() regardless of which run() that resolves to. + """ + + class ReportingWorker(threading.Thread): + def run(self): + tune.report(metric=123.0) + + def eval_subclassed_thread(config): + t = ReportingWorker() + t.start() + t.join(timeout=5) + return None + + analysis = tune.run( + eval_subclassed_thread, + config={"x": tune.uniform(0, 1)}, + metric="metric", + mode="min", + num_samples=1, + verbose=0, + ) + assert analysis.trials[0].last_result is not None, ( + "report() from a Thread SUBCLASS overriding run() was dropped; expected " + "_install_thread_context_propagation() to wrap _bootstrap_inner, not run(), " + "so an overridden run() is still covered" + ) + assert analysis.trials[0].last_result.get("metric") == 123.0 + + +def test_tune_log_records_from_worker_thread_reach_run_log(tmp_path): + """Follow-up to #996, third review point 4: automatic context + propagation (a plain worker Thread, or a ThreadPoolExecutor task) + carried report()'s runner/trial routing, but not `_state.log_run_id`, + which `_RunScopedFilter` matches against to decide whether a log record + belongs in THIS run's log file. A worker thread's own _TuneState starts + with log_run_id=None, so its log records were silently filtered out of + the run log even though its report() call correctly reached the right + trial. + """ + log_path = str(tmp_path / "worker_thread.log") + + def eval_logging_worker(config): + def worker(): + logger.info("MARKER_FROM_WORKER_THREAD") + + t = threading.Thread(target=worker) + t.start() + t.join(timeout=5) + return {"metric": 1.0} + + tune.run( + eval_logging_worker, + config={"x": tune.uniform(0, 1)}, + metric="metric", + mode="min", + num_samples=1, + verbose=2, + log_file_name=log_path, + ) + + text = open(log_path).read() + assert "MARKER_FROM_WORKER_THREAD" in text, ( + "a worker thread's log record was filtered out of its own run's log file; " + f"expected log_run_id propagation to carry it through, got: {text!r}" + ) + + +def test_late_report_via_stale_context_does_not_corrupt_finished_trial(): + """Follow-up to #996, fourth review point 2: a _RunContext captured for + a trial stays usable after that trial finishes. A worker thread that + captured get_run_context() and reports late, after tune.run() has + already returned and the trial is TERMINATED, used to still write + through: process_trial_result() overwrote the trial's already-final + metric_analysis/last_result with the late value, and report()'s own + trailing `if trial.is_finished(): raise StopIteration` (the normal + scheduler-stop signal for the CURRENT report) then raised into the late + caller too, which has no reason to expect it the way a trainable's own + control-flow loop does. + """ + captured = {} + + def eval_capturing(config): + captured["ctx"] = tune.get_run_context() + return {"metric": 1.0} + + analysis = tune.run( + eval_capturing, + config={"x": tune.uniform(0, 1)}, + metric="metric", + mode="min", + num_samples=1, + verbose=0, + ) + trial = analysis.trials[0] + assert trial.is_finished(), "expected the trial to be TERMINATED once tune.run() returns" + last_result_before = dict(trial.last_result) + metric_analysis_before = {k: dict(v) for k, v in trial.metric_analysis.items()} + + with tune.use_run_context(captured["ctx"]): + tune.report(metric=999.0) + + assert ( + trial.last_result == last_result_before + ), f"a late report corrupted the finished trial's last_result: {trial.last_result}" + assert ( + trial.metric_analysis == metric_analysis_before + ), f"a late report corrupted the finished trial's metric_analysis: {trial.metric_analysis}" + + +def test_tune_log_records_from_executor_worker_reach_run_log(tmp_path): + """Same as above (third review point 4), for a task submitted to a + ThreadPoolExecutor rather than a plain Thread, since the two are + separate propagation paths (_install_thread_context_propagation() vs + _install_executor_context_propagation()). + """ + log_path = str(tmp_path / "executor_worker.log") + executor = concurrent.futures.ThreadPoolExecutor(max_workers=1) + executor.submit(lambda: None).result() # warm up before any run exists + + def emit(): + logger.info("MARKER_FROM_EXECUTOR_WORKER") + + def eval_logging_executor(config): + future = executor.submit(emit) + future.result(timeout=5) + return {"metric": 1.0} + + try: + tune.run( + eval_logging_executor, + config={"x": tune.uniform(0, 1)}, + metric="metric", + mode="min", + num_samples=1, + verbose=2, + log_file_name=log_path, + ) + finally: + executor.shutdown(wait=True) + + text = open(log_path).read() + assert "MARKER_FROM_EXECUTOR_WORKER" in text, ( + "an executor worker's log record was filtered out of its own run's log file; " + f"expected log_run_id propagation to carry it through, got: {text!r}" + ) + + +def test_concurrent_reports_for_same_trial_admit_atomically(): + """Follow-up to #996, fifth review point 2: trial.is_finished() and + runner.process_trial_result() were not atomic. Two worker threads + reporting for the SAME trial through a shared propagated context (the + same handoff every test above uses, just two of them at once instead of + one) could both pass the is_finished() check while the trial was still + running; whichever one's process_trial_result() call landed AFTER the + other one's scheduler decision had already finished the trial still + wrote through, silently replacing the trial's real final result with a + stale one. + + Forced deterministically, with events rather than a sleep: worker B is + paused right after its own is_finished() check returns False, before it + enters the critical section, until worker A's entire report() call, + including the scheduler decision that terminates the trial, has + completed. B is released only then, so it always reaches admission with + an already-finished trial, which is exactly the window the fix closes + with a second, lock-protected check immediately before + process_trial_result(). + + Gate point: `_admission_lock_for(trial)`, patched to pause AFTER + fetching the real per-trial lock but BEFORE returning it, i.e. before + the `with` statement that follows ever calls `.acquire()` on it, so B + is paused without holding the lock. This used to gate on + `_next_training_iteration` instead (also called between the fast check + and the `with` block, at the time); #996 follow-up sixth review point 3 + moved that call INSIDE the locked section (allocation and admission + must be one atomic step, not two), so gating there now would mean B + pauses WHILE HOLDING the lock A's own report() needs, deadlocking both + threads instead of testing anything. + """ + + class StopOnFirstResult: + """A minimal TrialScheduler, the same pluggable interface + SequentialTrialRunner feeds any real scheduler through: STOPs the + trial the first time a result is admitted, so whichever worker + reaches admission first legitimately finishes the trial. + """ + + def set_search_properties(self, metric=None, mode=None, **spec): + pass + + def on_trial_add(self, runner, trial): + pass + + def on_trial_result(self, runner, trial, result): + return "STOP" + + def on_trial_complete(self, runner, trial_id, result=None, error=False): + pass + + def on_trial_remove(self, runner, trial): + pass + + b_ready = threading.Event() + a_finished = threading.Event() + orig_admission_lock_for = tune.tune._admission_lock_for + + def gated_admission_lock_for(trial_obj): + lock = orig_admission_lock_for(trial_obj) + if threading.current_thread().name == "B_WORKER": + b_ready.set() + assert a_finished.wait(timeout=5), "worker A never finished while B was paused" + return lock + + def eval_concurrent_reporters(config): + ctx = tune.get_run_context() + + def worker_a(): + try: + with tune.use_run_context(ctx): + tune.report(metric=1.0) + except StopIteration: + pass + finally: + a_finished.set() + + def worker_b(): + try: + with tune.use_run_context(ctx): + tune.report(metric=2.0) + except StopIteration: + pass + + thread_b = threading.Thread(target=worker_b, name="B_WORKER") + thread_b.start() + assert b_ready.wait(timeout=5), "worker B never reached the admission gap" + thread_a = threading.Thread(target=worker_a, name="A_WORKER") + thread_a.start() + thread_a.join(timeout=5) + thread_b.join(timeout=5) + return None + + with mock.patch("flaml.tune.tune._admission_lock_for", side_effect=gated_admission_lock_for): + analysis = tune.run( + eval_concurrent_reporters, + config={"x": tune.uniform(0, 1)}, + metric="metric", + mode="min", + num_samples=1, + scheduler=StopOnFirstResult(), + verbose=0, + ) + + trial = analysis.trials[0] + assert trial.is_finished() + assert trial.last_result.get("metric") == 1.0, ( + "worker B's report, paused mid-flight and released only after worker A's report " + "had already terminated the trial via the scheduler's STOP decision, still " + f"overwrote the trial's real final result; got {trial.last_result}" + ) + + +def test_tune_report_from_prestarted_reused_queue_worker_via_explicit_context(): + """Follow-up to #996, fifth review point 1 (restating fourth review + point 1): a generic persistent queue-consumer thread started BEFORE any + tune.run() call exists has nothing for _install_thread_context_ + propagation() to capture at its Thread.start() time, and it is not a + ThreadPoolExecutor either, so _install_executor_context_propagation() + does not apply. + + Automatic per-dispatch propagation for an arbitrary queue worker is not + being added: report()'s only handle on which trial a dequeued item + belongs to is whatever context was captured when it was produced, and + nothing at report()-call time can reconstruct that after the fact. + SequentialTrialRunner.step() reassigns runner.running_trial to a new + trial every step, so a fallback that resolved the trial live at + report() time would attribute a delayed item to whichever trial happens + to be running when it is finally dequeued, not the one it was produced + for, which is a silent wrong-trial write and strictly worse than today. + + The compatibility route the review also names already exists: + get_run_context()/use_run_context(), captured by the producer at + enqueue time and carried on the queue item itself instead of resolved + at dequeue time. This is that route, on exactly the shape described: a + worker thread started before any tune.run() call exists, reused + unmodified across two separate trials. + """ + work_queue = queue.Queue() + stop = object() + + def worker(): + while True: + item = work_queue.get() + try: + if item is stop: + return + ctx, value = item + with tune.use_run_context(ctx): + tune.report(metric=value) + finally: + work_queue.task_done() + + worker_thread = threading.Thread(target=worker) + worker_thread.start() # started before any tune.run() call exists + + def eval_via_prestarted_queue(config): + ctx = tune.get_run_context() + work_queue.put((ctx, config["x"])) + work_queue.join() # wait for the pre-started worker to drain this trial's item + return None + + try: + analysis = tune.run( + eval_via_prestarted_queue, + config={"x": tune.uniform(0, 1)}, + points_to_evaluate=[{"x": 11.0}, {"x": 22.0}], + metric="metric", + mode="min", + num_samples=2, + verbose=0, + ) + finally: + work_queue.put(stop) + worker_thread.join(timeout=5) + + reported = sorted(t.last_result.get("metric") for t in analysis.trials if t.last_result is not None) + assert reported == [ + 11.0, + 22.0, + ], ( + "a worker thread started before tune.run() and reused across both trials, using " + "the documented get_run_context()/use_run_context() API, did not route each " + f"trial's report to its own trial; got {reported}" + ) + + +def test_stop_trial_shares_lifecycle_lock_with_late_report(): + """Follow-up to #996, sixth review point 1: report()'s admission + (process_trial_result(), guarded by tune.py's per-trial + _admission_lock_for) and trial_runner.py's stop_trial() did not share a + lock, even though both mutate the same trial: last_result and + metric_analysis (via Trial.update_last_result()), status, and the + search_alg/scheduler on_trial_* callbacks. A trainable can report from + a background thread it does not wait for (every worker-thread test + above uses exactly this shape; here evaluation_function() just does not + join() it before returning), and run()'s own sequential loop calls + runner.stop_trial(trial_to_run) immediately after evaluation_function() + returns, whether or not that straggling report has finished. + + Forced deterministically, with events rather than a sleep: + Trial.update_last_result() (called from inside process_trial_result(), + itself inside the admission lock) is patched to signal it has been + entered and then block. Once that signal arrives, stop_trial() is + called directly from a separate thread, not through run()'s own loop: + doing it through the loop would block the driving thread on the very + lock this test needs to release the worker to unblock, deadlocking the + test itself rather than exercising the race. Before the fix, + stop_trial() had nothing to wait on: it read trial.last_result while + the straggling report's update_last_result() was still paused mid-write + (last_result still the trial's pre-report default), and handed that + stale value to the search algorithm's on_trial_complete() as the + trial's supposedly final result. After the fix, stop_trial() blocks on + the same per-trial lock until the report's entire critical section has + completed, so it always sees the trial's real last_result. + """ + worker_in_update = threading.Event() + release_worker = threading.Event() + orig_update_last_result = Trial.update_last_result + + def gated_update_last_result(self, result): + worker_in_update.set() + assert release_worker.wait(timeout=5), "test never released the in-flight report" + return orig_update_last_result(self, result) + + class RecordingSearchAlg: + """Minimal search_alg double: suggests exactly one trial, records + every result stop_trial() hands its on_trial_complete(). + """ + + def __init__(self): + self._suggested = False + self.complete_calls = [] + + def set_search_properties(self, metric=None, mode=None, config=None, **spec): + return True + + def suggest(self, trial_id): + if self._suggested: + return None + self._suggested = True + return {"x": 0.5} + + def on_trial_result(self, trial_id, result): + pass + + def on_trial_complete(self, trial_id, result=None, error=False): + self.complete_calls.append(dict(result) if result else result) + + search_alg = RecordingSearchAlg() + worker_errors = [] + + def eval_fire_and_forget(config): + ctx = tune.get_run_context() + trial = ctx.running_trial + runner = ctx.runner + + def worker(): + try: + with tune.use_run_context(ctx): + tune.report(metric=1.0) + except StopIteration: + pass + except Exception as exc: # pragma: no cover - failure diagnostics only + worker_errors.append(exc) + + report_thread = threading.Thread(target=worker, name="STRAGGLER") + report_thread.start() + assert worker_in_update.wait(timeout=5), "worker never reached update_last_result" + + stop_started = threading.Event() + + def call_stop_trial(): + stop_started.set() + runner.stop_trial(trial) + + stop_thread = threading.Thread(target=call_stop_trial, name="STOPPER") + stop_thread.start() + assert stop_started.wait(timeout=5), "stop_trial() thread never started" + release_worker.set() + report_thread.join(timeout=5) + stop_thread.join(timeout=5) + return None # run()'s own post-eval stop_trial() call becomes a harmless no-op + + with mock.patch.object(Trial, "update_last_result", gated_update_last_result): + analysis = tune.run( + eval_fire_and_forget, + config={"x": tune.uniform(0, 1)}, + metric="metric", + mode="min", + num_samples=1, + search_alg=search_alg, + verbose=0, + ) + + assert not worker_errors, f"straggling report thread raised: {worker_errors}" + trial = analysis.trials[0] + assert trial.last_result is not None and trial.last_result.get("metric") == 1.0 + assert search_alg.complete_calls, "stop_trial() never called on_trial_complete" + assert search_alg.complete_calls[0] is not None and search_alg.complete_calls[0].get("metric") == 1.0, ( + "stop_trial() handed the search algorithm a stale/incomplete last_result while the " + f"straggling report's update_last_result() was still in flight; got {search_alg.complete_calls[0]}" + ) + + +def test_training_iteration_allocated_inside_admission_lock(): + """Follow-up to #996, sixth review point 3: training_iteration used to + be allocated (_next_training_iteration()) BEFORE the per-trial + admission lock (_admission_lock_for) was acquired, not inside it. Two + truly concurrent reports for the same trial could then be handed + iteration numbers in one order and reach runner.process_trial_result(), + which is what a scheduler/searcher that orders trials by + training_iteration (ASHA and similar) actually sees, in the OTHER + order, if the thread that allocated the LATER number happened to reach + the lock first. + + Rather than trying to force that exact reordering (which the fix makes + impossible to construct at all, since allocation now only happens while + already holding the lock), this proves the mechanism directly: a + report's iteration allocation is patched to pause mid-call, and a + second, independent attempt to enter the SAME trial's admission section + is made concurrently. If allocation and admission share one critical + section, that second attempt must block for as long as allocation is + paused; if they are two separate critical sections (the bug), the + second attempt sails through immediately, since nothing is held during + allocation. + """ + allocating = threading.Event() + release_allocation = threading.Event() + orig_next_iteration = tune.tune._next_training_iteration + + def gated_next_iteration(trial_obj): + allocating.set() + assert release_allocation.wait(timeout=5), "test never released the paused allocation" + return orig_next_iteration(trial_obj) + + def eval_probe(config): + ctx = tune.get_run_context() + trial = ctx.running_trial + + def worker(): + with tune.use_run_context(ctx): + tune.report(metric=1.0) + + report_thread = threading.Thread(target=worker, name="ALLOCATOR") + report_thread.start() + assert allocating.wait(timeout=5), "worker never reached iteration allocation" + + probe_acquired = threading.Event() + + def probe(): + with tune.tune._admission_lock_for(trial): + probe_acquired.set() + + probe_thread = threading.Thread(target=probe, name="PROBE") + probe_thread.start() + acquired_while_allocation_paused = probe_acquired.wait(timeout=0.5) + + release_allocation.set() + report_thread.join(timeout=5) + probe_thread.join(timeout=5) + + assert not acquired_while_allocation_paused, ( + "a second, independent attempt to admit a report for the same trial acquired " + "the admission lock while training_iteration allocation for an in-flight report " + "was still paused: allocation is not happening inside the same critical section " + "as process_trial_result(), so a concurrent report can be allocated a later " + "iteration and still be admitted first" + ) + return None + + with mock.patch("flaml.tune.tune._next_training_iteration", side_effect=gated_next_iteration): + tune.run( + eval_probe, + config={"x": tune.uniform(0, 1)}, + metric="metric", + mode="min", + num_samples=1, + verbose=0, + ) + + +def test_logger_level_restored_as_inherited_not_explicit(): + """Follow-up to #996, sixth review point 4: _logger_level_enter() saved + logger.getEffectiveLevel(), the RESOLVED level after walking up the + logger hierarchy when this logger has no level of its own (logger.level + == logging.NOTSET), and _logger_level_exit() restored that resolved + number via logger.setLevel(). That gives the logger a permanent + explicit level it never had before: a logger that was inheriting must + go back to inheriting once every active run has exited, not end up + pinned to whatever the ancestor chain resolved to during the run. + + Not a contrived setup: flaml/__init__.py sets the "flaml" logger to + INFO at import time, so flaml.tune.logger's own getEffectiveLevel() + already resolves to INFO (via that ancestor) in any unconfigured + process. No level manipulation beyond resetting flaml.tune.logger's own + level to NOTSET is needed to construct "this logger is inheriting", the + precondition the fix is about; the assertion just below confirms it. + """ + saved_level = logger.level + try: + logger.setLevel(logging.NOTSET) # construct the precondition: inheriting, not explicit + assert logger.getEffectiveLevel() != logging.NOTSET, ( + "test assumes an ancestor logger (flaml/__init__.py sets 'flaml' to INFO) resolves " + "to a non-NOTSET effective level here" + ) + + tune.run( + lambda config: {"metric": 1.0}, + config={"x": tune.uniform(0, 1)}, + metric="metric", + mode="min", + num_samples=1, + verbose=1, + ) + + assert logger.level == logging.NOTSET, ( + "tune.run() left flaml.tune's logger pinned to an explicit level instead of " + f"restoring it to NOTSET (inheriting); got {logging.getLevelName(logger.level)}" + ) + finally: + logger.setLevel(saved_level) + + +def test_tune_report_from_unchanged_legacy_queue_worker_preserves_reporting_when_unambiguous(tmp_path): + """Follow-up to #996, eighth review point 1: a generic queue/callback + worker started before any tune.run() call exists, that never calls + get_run_context()/use_run_context() itself (an unchanged legacy + caller), used to have its tune.report() unconditionally dropped + (seventh review), with only a one-time warning to show for it. The + eighth review objected to that: the new regression test for it + asserted the drop (`last_result is None`) as the goal, codifying + backward-incompatible data loss instead of preserving the caller that + genuinely worked before #996. The pre-#996 module-global design + attributed exactly this caller shape correctly whenever only one + tune.run() call was active at a time, which is also the common real + case for this pattern (one long-lived worker draining a queue one run + at a time). This restores that: when exactly one sequential + tune.run() call is active anywhere in the process, + _resolve_unambiguous_active_runner lets a context-less report() + resolve to it automatically, with no change on the caller's part. + (Two or more concurrently active runs is a different, genuinely + ambiguous case; see + test_context_less_report_raises_when_multiple_runs_are_concurrently_active + below.) + + Reused, unmodified, across TWO separate trials (points_to_evaluate), + to prove this is a real per-dispatch resolution, landing each trial's + own report on that trial, not an accident of only ever having tried it + once. + """ + work_queue = queue.Queue() + stop = object() + + def worker(): + while True: + item = work_queue.get() + try: + if item is stop: + return + # unchanged legacy caller: no get_run_context()/use_run_context(), + # exactly the shape the review describes. + tune.report(metric=item) + finally: + work_queue.task_done() + + def eval_via_legacy_queue(config): + work_queue.put(config["x"]) + work_queue.join() # wait for the pre-started worker to drain this trial's item + return None + + log_path = str(tmp_path / "legacy_worker.log") + worker_thread = threading.Thread(target=worker) + try: + worker_thread.start() # started before any tune.run() call exists + analysis = tune.run( + eval_via_legacy_queue, + config={"x": tune.uniform(0, 1)}, + points_to_evaluate=[{"x": 11.0}, {"x": 22.0}], + metric="metric", + mode="min", + num_samples=2, + verbose=1, + log_file_name=log_path, + ) + finally: + # Put `stop` and join unconditionally: the worker thread is a plain + # non-daemon Thread blocked on queue.get() with nothing else able + # to release it, and leaving it running hangs the whole + # interpreter at process exit. + work_queue.put(stop) + worker_thread.join(timeout=5) + + reported = sorted(t.last_result.get("metric") for t in analysis.trials if t.last_result is not None) + assert reported == [ + 11.0, + 22.0, + ], ( + "an unchanged legacy queue worker, with the single active run this process had " + f"at the time, should have had both trials' reports preserved; got {reported}" + ) + + log_text = open(log_path).read() + assert "no active run to attribute it to" not in log_text, ( + "the context-less-report warning fired even though the report was successfully " + f"attributed and preserved: {log_text!r}" + ) + + +def test_context_less_report_raises_when_multiple_runs_are_concurrently_active(): + """Follow-up to #996, eighth review point 1: the single-active-run + fallback above is only safe because there is exactly one run it could + mean. With two tune.run() calls concurrently active, a context-less + report (no thread-local runner, no propagated/explicit context) + cannot be attributed to either one safely: guessing would silently + write one run's metric onto a different run's trial, worse than + dropping it. _resolve_unambiguous_active_runner raises RuntimeError + in this case instead. + + Forced deterministically, same pause/release shape as + test_concurrent_tune_run_does_not_corrupt_state: both driving threads + are paused inside their own evaluation function, so both runs are + genuinely active at once, then a THIRD, independent thread with no + context of its own (started from the test thread, never from inside + either run's evaluation_function(), so it inherits no propagated + context from either) calls tune.report() while both are paused. + """ + a_paused = threading.Event() + b_paused = threading.Event() + release_a = threading.Event() + release_b = threading.Event() + + def eval_a(config): + a_paused.set() + assert release_a.wait(timeout=5), "thread A was never released" + return {"metric": 1.0} + + def eval_b(config): + b_paused.set() + assert release_b.wait(timeout=5), "thread B was never released" + return {"metric": 2.0} + + def run_a(): + tune.run(eval_a, config={"x": tune.uniform(0, 1)}, metric="metric", mode="min", num_samples=1, verbose=0) + + def run_b(): + tune.run(eval_b, config={"x": tune.uniform(0, 1)}, metric="metric", mode="min", num_samples=1, verbose=0) + + thread_a = threading.Thread(target=run_a) + thread_b = threading.Thread(target=run_b) + thread_a.start() + assert a_paused.wait(timeout=5), "thread A never reached its evaluation function" + thread_b.start() + assert b_paused.wait(timeout=5), "thread B never reached its evaluation function" + + outcome = {} + + def context_less_reporter(): + try: + tune.report(metric=99.0) + except Exception as e: # noqa: BLE001 (captured for the assertion below) + outcome["result"] = ("error", e) + else: + outcome["result"] = ("ok", None) + + try: + reporter_thread = threading.Thread(target=context_less_reporter) + reporter_thread.start() + reporter_thread.join(timeout=5) + finally: + release_a.set() + release_b.set() + thread_a.join(timeout=5) + thread_b.join(timeout=5) + + status, err = outcome.get("result", (None, None)) + assert status == "error", ( + f"a context-less report with two runs concurrently active should have raised " + f"instead of silently dropping or guessing; got {outcome.get('result')}" + ) + assert isinstance(err, RuntimeError), f"expected RuntimeError, got {type(err).__name__}: {err}" + + +def test_context_less_report_warns_and_drops_when_no_run_is_active(caplog): + """Follow-up to #996, eighth review points 1 and 3: with no tune.run() + call active anywhere in the process, a context-less report has + nothing to attribute to even by inference (contrast the + single-active-run fallback above), and is still dropped, with a + one-time-per-run warning. The warning now goes through a dedicated + logger, `flaml.tune.diagnostics`, never a run-scoped one (eighth + review point 3), so it is captured here the way it actually surfaces: + propagated up to the root logger, not read out of any run's own log + file (there is no active run in this test for one to belong to). + """ + with mock.patch.object(tune_module, "_context_less_report_warned", False): + with caplog.at_level(logging.WARNING, logger="flaml.tune.diagnostics"): + result = tune.report(metric=1.0) + assert result is None, "a context-less report with no active run should be dropped, not raise or record" + assert ( + "no active run to attribute it to" in caplog.text + ), f"expected the context-less-report warning, got: {caplog.text!r}" + + +def test_diagnostics_warning_never_reaches_a_run_owned_log_file(tmp_path): + """Follow-up to #996, eighth review point 3: the `flaml_tune_unscoped` + bypass this replaced sent an unattributed warning to EVERY + concurrently active run's own log handler (via _RunScopedFilter), so + an unrelated diagnostic from one caller could land inside a DIFFERENT + run's log file. `_diagnostics_logger` (used by the context-less-report + warning) replaces that: verified directly against the logging + mechanism itself, not just today's call sites, so a future caller of + `_diagnostics_logger` inherits the same guarantee. Even with a real + run active and its own FileHandler attached, a record emitted on + `_diagnostics_logger` does not appear in that run's log file. + """ + log_path = str(tmp_path / "run.log") + seen = threading.Event() + + def eval_and_emit(config): + tune_module._diagnostics_logger.warning("MARKER_DIAGNOSTICS_ONLY") + seen.set() + return {"metric": 1.0} + + tune.run( + eval_and_emit, + config={"x": tune.uniform(0, 1)}, + metric="metric", + mode="min", + num_samples=1, + verbose=2, + log_file_name=log_path, + ) + assert seen.is_set(), "the evaluation function never ran" + log_text = open(log_path).read() + assert ( + "MARKER_DIAGNOSTICS_ONLY" not in log_text + ), f"a _diagnostics_logger record leaked into a run-owned log file: {log_text!r}"