From 1eb58febda389e77b42913895cde6b107d7bd0f9 Mon Sep 17 00:00:00 2001 From: Amir Fathi Date: Tue, 22 Sep 2026 23:54:37 +0000 Subject: [PATCH 01/10] fix: tune.run() from concurrent threads corrupts shared runner state report()/run() coordinated through five module-level globals (_use_ray, _runner, _verbose, _running_trial, _training_iteration), shared by every thread. Two threads calling tune.run() concurrently overwrite each other's runner/trial bookkeeping mid-flight: whichever thread's run() call most recently assigns _runner wins the global for every thread's subsequent report()/stop_trial() calls, including resetting it while another thread is still running. Move that state into a threading.local subclass so each thread gets its own copy. The existing single-thread nested-reentrancy save/restore in run()'s try/finally is unchanged, just scoped per thread instead of process-wide. Fixes #996 --- flaml/tune/tune.py | 126 ++++++++++++++++--------------- test/tune/test_concurrent_run.py | 102 +++++++++++++++++++++++++ 2 files changed, 169 insertions(+), 59 deletions(-) create mode 100644 test/tune/test_concurrent_run.py diff --git a/flaml/tune/tune.py b/flaml/tune/tune.py index ea2abfe2a6..bcce89249a 100644 --- a/flaml/tune/tune.py +++ b/flaml/tune/tune.py @@ -5,6 +5,7 @@ import datetime import os import sys +import threading import time from collections import defaultdict from typing import Callable, Dict, List, Optional, Tuple, Union @@ -41,11 +42,27 @@ internal_mlflow = False -_use_ray = True -_runner = None -_verbose = 0 -_running_trial = None -_training_iteration = 0 +class _TuneState(threading.local): + """Per-thread run() state. + + report()/run() coordinate purely through this state (use_ray, runner, + verbose, running_trial, training_iteration). 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. + """ + + def __init__(self): + self.use_ray = True + self.runner = None + self.verbose = 0 + self.running_trial = None + self.training_iteration = 0 + + +_state = _TuneState() INCUMBENT_RESULT = "__incumbent_result__" @@ -189,11 +206,7 @@ def compute_with_config(config): 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: + if _state.use_ray: try: from ray import __version__ as ray_version @@ -211,22 +224,22 @@ def compute_with_config(config): result = kwargs if _metric is not None: result[DEFAULT_METRIC] = _metric - trial = getattr(_runner, "running_trial", None) + trial = getattr(_state.runner, "running_trial", None) if not trial: return None - if _running_trial == trial: - _training_iteration += 1 + if _state.running_trial == trial: + _state.training_iteration += 1 else: - _training_iteration = 0 - _running_trial = trial - result["training_iteration"] = _training_iteration + _state.training_iteration = 0 + _state.running_trial = trial + result["training_iteration"] = _state.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: + _state.runner.process_trial_result(trial, result) + if _state.verbose > 2: logger.info(f"result: {result}") if trial.is_finished(): raise StopIteration @@ -478,15 +491,11 @@ def easy_objective(config): **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 + old_use_ray = _state.use_ray + old_verbose = _state.verbose + old_running_trial = _state.running_trial + old_training_iteration = _state.training_iteration if log_file_name: dir_name = os.path.dirname(log_file_name) @@ -498,13 +507,12 @@ def easy_objective(config): 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 + _state.use_ray = False + _state.verbose = verbose old_handlers = logger.handlers old_level = logger.getEffectiveLevel() logger.handlers = [] - global _runner - old_runner = _runner + old_runner = _state.runner assert not ray_args, "ray_args is only valid when use_ray=True" if ( old_handlers @@ -674,7 +682,7 @@ def easy_objective(config): 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 + _state.use_ray = True try: analysis = tune.run( evaluation_function, @@ -695,10 +703,10 @@ def easy_objective(config): 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 + _state.use_ray = old_use_ray + _state.verbose = old_verbose + _state.running_trial = old_running_trial + _state.training_iteration = old_training_iteration if use_spark: # parallel run with spark @@ -759,7 +767,7 @@ def easy_objective(config): with parallel_backend("spark"): with Parallel(n_jobs=n_concurrent_trials, verbose=max(0, (verbose - 1) * 50)) as parallel: try: - _runner = SparkTrialRunner( + _state.runner = SparkTrialRunner( search_alg=search_alg, scheduler=scheduler, metric=metric, @@ -779,9 +787,9 @@ def easy_objective(config): 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: + while len(_state.runner.running_trials) < n_concurrent_trials: # suggest trials for spark - trial_next = _runner.step() + trial_next = _state.runner.step() if trial_next: num_trials += 1 else: @@ -789,13 +797,13 @@ def easy_objective(config): logger.debug(f"consecutive failures is {num_failures}") if num_failures >= upperbound_num_failures: break - trials_to_run = _runner.running_trials + 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(_runner.running_trials)} RUNNING," - f" {len(_runner._trials) - len(_runner.running_trials)} TERMINATED" + 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]}" @@ -817,7 +825,7 @@ def easy_objective(config): while results: result = results.pop(0) trial_to_run = trials_to_run[0] - _runner.running_trial = trial_to_run + _state.runner.running_trial = trial_to_run if result is not None: if _internal_mlflow: mlflow_integration.record_trial(result, trial_to_run, metric) @@ -832,10 +840,10 @@ def easy_objective(config): else: logger.info("Brief result: {metric: result}") report(_metric=result) - _runner.stop_trial(trial_to_run) + _state.runner.stop_trial(trial_to_run) num_failures = 0 analysis = ExperimentAnalysis( - _runner.get_trials(), + _state.runner.get_trials(), metric=metric, mode=mode, lexico_objectives=lexico_objectives, @@ -857,12 +865,12 @@ def easy_objective(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 + _state.use_ray = old_use_ray + _state.verbose = old_verbose + _state.running_trial = old_running_trial + _state.training_iteration = old_training_iteration if not use_ray: - _runner = old_runner + _state.runner = old_runner logger.handlers = old_handlers logger.setLevel(old_level) if _internal_mlflow: @@ -870,13 +878,13 @@ def easy_objective(config): # simple sequential run without using tune.run() from ray time_start = time.time() - _use_ray = False + _state.use_ray = False if scheduler: scheduler.set_search_properties(metric=metric, mode=mode) from .trial_runner import SequentialTrialRunner try: - _runner = SequentialTrialRunner( + _state.runner = SequentialTrialRunner( search_alg=search_alg, scheduler=scheduler, metric=metric, @@ -892,7 +900,7 @@ def easy_objective(config): and (num_samples < 0 or num_trials < num_samples) and num_failures < upperbound_num_failures ): - trial_to_run = _runner.step() + trial_to_run = _state.runner.step() if trial_to_run: num_trials += 1 if verbose: @@ -913,7 +921,7 @@ def easy_objective(config): trial_to_run.set_status(Trial.ERROR) else: report(_metric=result) - _runner.stop_trial(trial_to_run) + _state.runner.stop_trial(trial_to_run) num_failures = 0 if trial_to_run.last_result is None: # application stops tuning by returning None @@ -925,7 +933,7 @@ def easy_objective(config): 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(), + _state.runner.get_trials(), metric=metric, mode=mode, lexico_objectives=lexico_objectives, @@ -946,12 +954,12 @@ def easy_objective(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 + _state.use_ray = old_use_ray + _state.verbose = old_verbose + _state.running_trial = old_running_trial + _state.training_iteration = old_training_iteration if not use_ray: - _runner = old_runner + _state.runner = old_runner logger.handlers = old_handlers logger.setLevel(old_level) if _internal_mlflow: diff --git a/test/tune/test_concurrent_run.py b/test/tune/test_concurrent_run.py new file mode 100644 index 0000000000..af4c3579db --- /dev/null +++ b/test/tune/test_concurrent_run.py @@ -0,0 +1,102 @@ +"""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 threading + +from flaml import tune + + +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}" From 7c1913481f81de0872cb4144e9b0431575d75c91 Mon Sep 17 00:00:00 2001 From: Amir Fathi Date: Wed, 23 Sep 2026 03:26:07 +0000 Subject: [PATCH 02/10] fix: address CHANGES_REQUESTED review on #996 thread-safety fix Scope the shared flaml.tune.logger handlers/level per run instead of replacing logger.handlers wholesale, so concurrent tune.run() calls no longer lose or leak each other's log routing (the logger was a sixth piece of process-global state left over after the five _TuneState fields were made thread-local). Wrap tune.run()'s setup (searcher/scheduler construction) in the same restore path used on normal return, so a setup failure no longer skips the state/logger restore and leaks mutated state into the next call on that thread. Add get_run_context()/use_run_context() so a trainable that spawns its own worker/callback thread can hand that thread the driving thread's run state before calling tune.report() from it, which the thread-local _TuneState broke silently. --- flaml/tune/__init__.py | 2 +- flaml/tune/tune.py | 2237 ++++++++++++++++-------------- test/tune/test_concurrent_run.py | 210 +++ 3 files changed, 1401 insertions(+), 1048 deletions(-) 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/tune.py b/flaml/tune/tune.py index bcce89249a..f7474b9a2f 100644 --- a/flaml/tune/tune.py +++ b/flaml/tune/tune.py @@ -1,1047 +1,1190 @@ -# ! -# * 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 threading -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 - - -class _TuneState(threading.local): - """Per-thread run() state. - - report()/run() coordinate purely through this state (use_ray, runner, - verbose, running_trial, training_iteration). 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. - """ - - def __init__(self): - self.use_ray = True - self.runner = None - self.verbose = 0 - self.running_trial = None - self.training_iteration = 0 - - -_state = _TuneState() - -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. - """ - if _state.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(_state.runner, "running_trial", None) - if not trial: - return None - if _state.running_trial == trial: - _state.training_iteration += 1 - else: - _state.training_iteration = 0 - _state.running_trial = trial - result["training_iteration"] = _state.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 - _state.runner.process_trial_result(trial, result) - if _state.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 internal_mlflow - old_use_ray = _state.use_ray - old_verbose = _state.verbose - old_running_trial = _state.running_trial - old_training_iteration = _state.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: - _state.use_ray = False - _state.verbose = verbose - old_handlers = logger.handlers - old_level = logger.getEffectiveLevel() - logger.handlers = [] - old_runner = _state.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") - _state.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: - _state.use_ray = old_use_ray - _state.verbose = old_verbose - _state.running_trial = old_running_trial - _state.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: - _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 - finally: - # recover the global variables in case of nested run - _state.use_ray = old_use_ray - _state.verbose = old_verbose - _state.running_trial = old_running_trial - _state.training_iteration = old_training_iteration - if not use_ray: - _state.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() - _state.use_ray = False - if scheduler: - scheduler.set_search_properties(metric=metric, mode=mode) - from .trial_runner import SequentialTrialRunner - - try: - _state.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 = _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 - 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) - _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 variables in case of nested run - _state.use_ray = old_use_ray - _state.verbose = old_verbose - _state.running_trial = old_running_trial - _state.training_iteration = old_training_iteration - if not use_ray: - _state.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 contextlib +import datetime +import os +import sys +import threading +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 + + +class _TuneState(threading.local): + """Per-thread run() state. + + report()/run() coordinate purely through this state (use_ray, runner, + verbose, running_trial, training_iteration, log_run_id). 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 (e.g. + a background worker that later calls tune.report()) starts from these + defaults too, with no runner attached; see get_run_context()/ + use_run_context() below for the supported way to hand that thread the + calling thread's state. + """ + + def __init__(self): + self.use_ray = True + self.runner = None + self.verbose = 0 + self.running_trial = None + self.training_iteration = 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. + """ + + 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: + global _logger_pristine_level + with _logger_state_lock: + if not _active_log_levels: + _logger_pristine_level = logger.getEffectiveLevel() + _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 thread's active tune.run() state (#996 follow-up). + + `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 is silently + dropped, same as calling tune.report() outside of tune.run() entirely. + Capture the driving thread's context with get_run_context() and attach it + on the worker thread with use_run_context() to report through it. + """ + + __slots__ = ("use_ray", "runner", "verbose", "running_trial", "training_iteration") + + def __init__(self, use_ray, runner, verbose, running_trial, training_iteration): + self.use_ray = use_ray + self.runner = runner + self.verbose = verbose + self.running_trial = running_trial + self.training_iteration = training_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. + """ + if _state.runner is None: + return None + return _RunContext(_state.use_ray, _state.runner, _state.verbose, _state.running_trial, _state.training_iteration) + + +@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". + + Note this only propagates report() bookkeeping (which trial, which + training_iteration); 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 + old_use_ray = _state.use_ray + old_runner = _state.runner + old_verbose = _state.verbose + old_running_trial = _state.running_trial + old_training_iteration = _state.training_iteration + _state.use_ray = ctx.use_ray + _state.runner = ctx.runner + _state.verbose = ctx.verbose + _state.running_trial = ctx.running_trial + _state.training_iteration = ctx.training_iteration + try: + yield + finally: + _state.use_ray = old_use_ray + _state.runner = old_runner + _state.verbose = old_verbose + _state.running_trial = old_running_trial + _state.training_iteration = old_training_iteration + + +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. + """ + if _state.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(_state.runner, "running_trial", None) + if not trial: + return None + if _state.running_trial == trial: + _state.training_iteration += 1 + else: + _state.training_iteration = 0 + _state.running_trial = trial + result["training_iteration"] = _state.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 + _state.runner.process_trial_result(trial, result) + if _state.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 internal_mlflow + old_use_ray = _state.use_ray + old_verbose = _state.verbose + old_running_trial = _state.running_trial + old_training_iteration = _state.training_iteration + old_runner = _state.runner + old_log_run_id = _state.log_run_id + _run_handler = None + _internal_mlflow = False + mlflow_integration = None + + def _restore_tune_state(): + """Undo every mutation this call made to shared/thread-local state. + + Called from the tail of every branch below AND from the except + clause wrapping the setup section, so a searcher/scheduler/backend + setup failure restores state exactly like a normal return does, + instead of leaking a mutated _state/logger into whatever run() this + thread resumes next (#996 follow-up). + """ + _state.use_ray = old_use_ray + _state.verbose = old_verbose + _state.running_trial = old_running_trial + _state.training_iteration = old_training_iteration + if not use_ray: + _state.runner = old_runner + _state.log_run_id = old_log_run_id + 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 + + 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 + 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: + _restore_tune_state() + + 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: + _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 + finally: + # recover the global/shared state in case of nested run + _restore_tune_state() + + # 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 + + try: + _state.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 = _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 + 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) + _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 + _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 index af4c3579db..cf5fc826fc 100644 --- a/test/tune/test_concurrent_run.py +++ b/test/tune/test_concurrent_run.py @@ -19,9 +19,13 @@ (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 threading +import pytest + from flaml import tune +from flaml.tune.logger import logger def test_concurrent_tune_run_does_not_corrupt_state(): @@ -100,3 +104,209 @@ def metric_of(key): 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: 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. + + First half is a positive control: without using get_run_context()/ + use_run_context(), report() from that worker thread is silently dropped, + the same way calling tune.report() outside of tune.run() is documented + to be a no-op. This is a real, pre-existing limitation this PR does not + claim to fix by itself. Second half shows the supported way to fix it: + capture the driving thread's context and attach it on the worker thread. + """ + + def eval_without_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(), if it lands, is the only result + + analysis = tune.run( + eval_without_propagation, + config={"x": tune.uniform(0, 1)}, + metric="metric", + mode="min", + num_samples=1, + verbose=0, + ) + assert analysis.trials[0].last_result is None, ( + "expected report() from an un-propagated worker thread to be silently dropped " + f"(pre-existing limitation); got {analysis.trials[0].last_result}" + ) + + def eval_with_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_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 From 272ba31e48ace1c884756e8d5ca32a00f9cddbb4 Mon Sep 17 00:00:00 2001 From: Amir Fathi Date: Wed, 23 Sep 2026 07:39:06 +0000 Subject: [PATCH 03/10] fix: address second CHANGES_REQUESTED review on #996 thread-safety fix report() from a trainable-spawned thread now works with no caller-side change: run() pins a contextvars context around each evaluation_function() call, and a patched threading.Thread inherits it into any thread the trainable spawns on its own. get_run_context()/use_run_context() stay as an explicit escape hatch for cases the patch can't reach. training_iteration moved off the per-thread/per-context snapshot onto a lock-guarded counter keyed by trial, so repeated handoffs for the same trial keep counting instead of restarting from a stale copy. Spark backend init and the sequential scheduler setup are now inside the same outer try/finally as the rest of run(), so a failure there restores _state/logger like every other setup failure does. --- flaml/tune/tune.py | 423 ++++++++++++++++++++----------- test/tune/test_concurrent_run.py | 142 +++++++++-- 2 files changed, 401 insertions(+), 164 deletions(-) diff --git a/flaml/tune/tune.py b/flaml/tune/tune.py index f7474b9a2f..ff533de7af 100644 --- a/flaml/tune/tune.py +++ b/flaml/tune/tune.py @@ -3,11 +3,13 @@ # * Licensed under the MIT License. See LICENSE file in the # * project root for license information. 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 @@ -46,27 +48,27 @@ class _TuneState(threading.local): """Per-thread run() state. - report()/run() coordinate purely through this state (use_ray, runner, - verbose, running_trial, training_iteration, log_run_id). 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 (e.g. - a background worker that later calls tune.report()) starts from these - defaults too, with no runner attached; see get_run_context()/ - use_run_context() below for the supported way to hand that thread the - calling thread's 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.running_trial = None - self.training_iteration = 0 self.log_run_id = None @@ -128,25 +130,134 @@ def _logger_level_exit(level: int) -> None: class _RunContext: - """Snapshot of one thread's active tune.run() state (#996 follow-up). + """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 is silently - dropped, same as calling tune.report() outside of tune.run() entirely. - Capture the driving thread's context with get_run_context() and attach it - on the worker thread with use_run_context() to report through it. + `_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). """ - __slots__ = ("use_ray", "runner", "verbose", "running_trial", "training_iteration") + __slots__ = ("use_ray", "runner", "verbose", "running_trial") - def __init__(self, use_ray, runner, verbose, running_trial, training_iteration): + def __init__(self, use_ray, runner, verbose, running_trial): self.use_ray = use_ray self.runner = runner self.verbose = verbose self.running_trial = running_trial - self.training_iteration = training_iteration + + +# 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() 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 threading.Thread inherit the calling thread's contextvars + Context, 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. + Patching Thread.start()/run() this way is the standard trick other + libraries (structlog, OpenTelemetry) use to give plain threads the same + inheritance asyncio gets for free: start() captures + contextvars.copy_context() on the CALLING thread (the one invoking + .start(), which is "inside" evaluation_function whenever the trainable + itself is the one spawning the helper thread), and the new OS thread's + run() executes inside that captured Context. A thread started outside + any tune.run() call captures a context with nothing bound in it and + behaves exactly as before. Idempotent: a second import/call is a no-op. + """ + if getattr(threading.Thread, "_flaml_tune_context_propagation", False): + return + _orig_start = threading.Thread.start + _orig_run = threading.Thread.run + + def _start(self, *args, **kwargs): + self._flaml_tune_ctx = contextvars.copy_context() + return _orig_start(self, *args, **kwargs) + + def _run(self, *args, **kwargs): + ctx = getattr(self, "_flaml_tune_ctx", None) + if ctx is None: + return _orig_run(self, *args, **kwargs) + return ctx.run(_orig_run, self, *args, **kwargs) + + threading.Thread.start = _start + threading.Thread.run = _run + threading.Thread._flaml_tune_context_propagation = True + + +_install_thread_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() + + +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"]: @@ -156,10 +267,15 @@ def get_run_context() -> Optional["_RunContext"]: 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 + already gets a context automatically, via _propagated_context and + _install_thread_context_propagation(). This (and use_run_context()) + stay as the explicit form for the cases that patch cannot reach. """ if _state.runner is None: return None - return _RunContext(_state.use_ray, _state.runner, _state.verbose, _state.running_trial, _state.training_iteration) + return _RunContext(_state.use_ray, _state.runner, _state.verbose, _state.runner.running_trial) @contextlib.contextmanager @@ -171,33 +287,20 @@ def use_run_context(ctx: Optional["_RunContext"]): `ctx=None` is accepted and is a no-op, so callers do not need to special case "this thread never got a context". - Note this only propagates report() bookkeeping (which trial, which - training_iteration); 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. + Note this only propagates report() bookkeeping (which runner, which + trial); 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 - old_use_ray = _state.use_ray - old_runner = _state.runner - old_verbose = _state.verbose - old_running_trial = _state.running_trial - old_training_iteration = _state.training_iteration - _state.use_ray = ctx.use_ray - _state.runner = ctx.runner - _state.verbose = ctx.verbose - _state.running_trial = ctx.running_trial - _state.training_iteration = ctx.training_iteration + token = _propagated_context.set(ctx) try: yield finally: - _state.use_ray = old_use_ray - _state.runner = old_runner - _state.verbose = old_verbose - _state.running_trial = old_running_trial - _state.training_iteration = old_training_iteration + _propagated_context.reset(token) class ExperimentAnalysis(EA): @@ -338,8 +441,25 @@ def compute_with_config(config): 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. """ - if _state.use_ray: + use_ray = _state.use_ray + runner = _state.runner + verbose = _state.verbose + running_trial = None + if runner is None: + ctx = _propagated_context.get() + if ctx is not None: + use_ray = ctx.use_ray + runner = ctx.runner + verbose = ctx.verbose + running_trial = ctx.running_trial + if use_ray: try: from ray import __version__ as ray_version @@ -357,22 +477,21 @@ def compute_with_config(config): result = kwargs if _metric is not None: result[DEFAULT_METRIC] = _metric - trial = getattr(_state.runner, "running_trial", None) + # 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: return None - if _state.running_trial == trial: - _state.training_iteration += 1 - else: - _state.training_iteration = 0 - _state.running_trial = trial - result["training_iteration"] = _state.training_iteration + result["training_iteration"] = _next_training_iteration(trial) 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 - _state.runner.process_trial_result(trial, result) - if _state.verbose > 2: + runner.process_trial_result(trial, result) + if verbose > 2: logger.info(f"result: {result}") if trial.is_finished(): raise StopIteration @@ -627,8 +746,6 @@ def easy_objective(config): global internal_mlflow old_use_ray = _state.use_ray old_verbose = _state.verbose - old_running_trial = _state.running_trial - old_training_iteration = _state.training_iteration old_runner = _state.runner old_log_run_id = _state.log_run_id _run_handler = None @@ -638,16 +755,18 @@ def easy_objective(config): def _restore_tune_state(): """Undo every mutation this call made to shared/thread-local state. - Called from the tail of every branch below AND from the except - clause wrapping the setup section, so a searcher/scheduler/backend - setup failure restores state exactly like a normal return does, - instead of leaking a mutated _state/logger into whatever run() this - thread resumes next (#996 follow-up). + 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 - _state.running_trial = old_running_trial - _state.training_iteration = old_training_iteration if not use_ray: _state.runner = old_runner _state.log_run_id = old_log_run_id @@ -841,13 +960,23 @@ def _restore_tune_state(): _restore_tune_state() raise - 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 - try: + # 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, @@ -866,68 +995,65 @@ def _restore_tune_state(): for trial in analysis.trials: f.write(f"result: {trial.last_result}\n") return analysis - finally: - _restore_tune_state() - - 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." + + 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, ) - with parallel_backend("spark"): - with Parallel(n_jobs=n_concurrent_trials, verbose=max(0, (verbose - 1) * 50)) as parallel: - try: + 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, @@ -1024,18 +1150,14 @@ def _restore_tune_state(): # ) return analysis - finally: - # recover the global/shared state in case of nested run - _restore_tune_state() - # 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 + # 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 - try: _state.runner = SequentialTrialRunner( search_alg=search_alg, scheduler=scheduler, @@ -1058,8 +1180,20 @@ def _restore_tune_state(): 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) + # Pin this evaluation call's runner/trial in _propagated_context + # so a worker thread evaluation_function spawns on its own can + # call tune.report() and land on the right trial with no + # trainable-side change (#996 follow-up, second review point 1). + # _install_thread_context_propagation() is what makes a plain + # threading.Thread started inside this call see it. + _prop_token = _propagated_context.set( + _RunContext(_state.use_ray, _state.runner, _state.verbose, trial_to_run) + ) + 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: @@ -1105,7 +1239,8 @@ def _restore_tune_state(): return analysis finally: - # recover the global/shared state in case of nested run + # 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() diff --git a/test/tune/test_concurrent_run.py b/test/tune/test_concurrent_run.py index cf5fc826fc..98db4dd103 100644 --- a/test/tune/test_concurrent_run.py +++ b/test/tune/test_concurrent_run.py @@ -252,43 +252,50 @@ def test_tune_run_setup_failure_restores_state(): def test_tune_report_from_trainable_spawned_thread(): - """Follow-up to #996, reviewer point 1: 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. - - First half is a positive control: without using get_run_context()/ - use_run_context(), report() from that worker thread is silently dropped, - the same way calling tune.report() outside of tune.run() is documented - to be a no-op. This is a real, pre-existing limitation this PR does not - claim to fix by itself. Second half shows the supported way to fix it: - capture the driving thread's context and attach it on the worker 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_without_propagation(config): + 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(), if it lands, is the only result + return None # the worker thread's report() is the only result analysis = tune.run( - eval_without_propagation, + 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 None, ( - "expected report() from an un-propagated worker thread to be silently dropped " - f"(pre-existing limitation); got {analysis.trials[0].last_result}" + 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_propagation(config): + def eval_with_explicit_propagation(config): ctx = tune.get_run_context() def worker(): @@ -301,7 +308,7 @@ def worker(): return None analysis2 = tune.run( - eval_with_propagation, + eval_with_explicit_propagation, config={"x": tune.uniform(0, 1)}, metric="metric", mode="min", @@ -310,3 +317,98 @@ def worker(): ) 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: PySpark not installed, 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. + + PySpark is not installed in this test environment, so check_spark() + deterministically returns unavailable; no real Spark cluster needed. + """ + 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(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 From ab260838a99645fa90e6393d1e823a4e4f365ee0 Mon Sep 17 00:00:00 2001 From: Amir Fathi Date: Wed, 23 Sep 2026 10:06:57 +0000 Subject: [PATCH 04/10] fix: test_tune_run_spark_setup_failure_restores_state fails on CI legs with real pyspark installed The test relied on PySpark being absent so check_spark() would return unavailable. ubuntu-latest 3.11/3.12/3.13 CI legs install real pyspark (3.5.1/4.0.1/4.1.0), so check_spark() succeeds there and the test never raises ImportError. Patch check_spark() directly to force the failure path regardless of whether pyspark is actually installed. --- test/tune/test_concurrent_run.py | 33 +++++++++++++++++++------------- 1 file changed, 20 insertions(+), 13 deletions(-) diff --git a/test/tune/test_concurrent_run.py b/test/tune/test_concurrent_run.py index 98db4dd103..fe8257e5df 100644 --- a/test/tune/test_concurrent_run.py +++ b/test/tune/test_concurrent_run.py @@ -21,6 +21,7 @@ """ import threading +from unittest import mock import pytest @@ -371,29 +372,35 @@ 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: PySpark not installed, the same failure a user + 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. - PySpark is not installed in this test environment, so check_spark() - deterministically returns unavailable; no real Spark cluster needed. + 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 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, - ) + 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" From 16773f5a23b8f59ae3e90631d320404ff572e092 Mon Sep 17 00:00:00 2001 From: Amir Fathi Date: Wed, 23 Sep 2026 11:44:58 +0000 Subject: [PATCH 05/10] fix: address third CHANGES_REQUESTED review on #996 thread-safety fix Three gaps in the automatic context-propagation patch, all real: 1. ThreadPoolExecutor worker threads call Thread.start() once, when the pool spins them up, not once per submitted task, so a task submitted to an already-warmed-up worker never saw a context captured at start() time. Fixed by also patching ThreadPoolExecutor.submit() to attach the context per task rather than per thread; verified this is not otherwise handled (a plain contextvars.ContextVar is not propagated into an already-running pool worker either). 2. The patch replaced threading.Thread.run directly, which a Thread subclass overriding its own run() shadows. Now wraps Thread._bootstrap_inner instead, which CPython calls internally and which invokes self.run() regardless of which run() that resolves to. 3. The patch captured contextvars.copy_context() (the full ambient Context), dragging every application ContextVar into every new thread process-wide. Now captures and restores only FLAML's own _RunContext value. Folded log_run_id into _RunContext and propagate it through both paths too: a worker thread's log records were being filtered out of the run log the same way an un-propagated report() used to be dropped, since the run-scoped log filter checks a separate thread-local the context patch was not setting. Five new regression tests: pre-warmed executor worker, pre-warmed executor reused across two trials, a Thread subclass overriding run(), and log records reaching the run log through both the plain-thread and executor paths. All five fail on the prior commit and pass with this fix; the six pre-existing concurrency tests are unaffected. --- flaml/tune/tune.py | 165 +++++++++++++++++++++------ test/tune/test_concurrent_run.py | 186 +++++++++++++++++++++++++++++++ 2 files changed, 316 insertions(+), 35 deletions(-) diff --git a/flaml/tune/tune.py b/flaml/tune/tune.py index ff533de7af..56b5c73430 100644 --- a/flaml/tune/tune.py +++ b/flaml/tune/tune.py @@ -2,6 +2,7 @@ # * 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 @@ -152,15 +153,22 @@ class _RunContext: 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") + __slots__ = ("use_ray", "runner", "verbose", "running_trial", "log_run_id") - def __init__(self, use_ray, runner, verbose, running_trial): + 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 @@ -172,7 +180,8 @@ def __init__(self, use_ray, runner, verbose, running_trial): # 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() below is what makes that second, +# _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( @@ -181,45 +190,114 @@ def __init__(self, use_ray, runner, verbose, running_trial): def _install_thread_context_propagation() -> None: - """Make threading.Thread inherit the calling thread's contextvars - Context, process-wide, once. + """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. - Patching Thread.start()/run() this way is the standard trick other - libraries (structlog, OpenTelemetry) use to give plain threads the same - inheritance asyncio gets for free: start() captures - contextvars.copy_context() on the CALLING thread (the one invoking - .start(), which is "inside" evaluation_function whenever the trainable - itself is the one spawning the helper thread), and the new OS thread's - run() executes inside that captured Context. A thread started outside - any tune.run() call captures a context with nothing bound in it and + + 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_run = threading.Thread.run + _orig_bootstrap_inner = threading.Thread._bootstrap_inner def _start(self, *args, **kwargs): - self._flaml_tune_ctx = contextvars.copy_context() + self._flaml_tune_ctx = _propagated_context.get() return _orig_start(self, *args, **kwargs) - def _run(self, *args, **kwargs): + def _bootstrap_inner(self, *args, **kwargs): ctx = getattr(self, "_flaml_tune_ctx", None) if ctx is None: - return _orig_run(self, *args, **kwargs) - return ctx.run(_orig_run, self, *args, **kwargs) + 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.run = _run + 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. @@ -268,14 +346,18 @@ def get_run_context() -> Optional["_RunContext"]: 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 - already gets a context automatically, via _propagated_context and - _install_thread_context_propagation(). This (and use_run_context()) - stay as the explicit form for the cases that patch cannot reach. + 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) + return _RunContext(_state.use_ray, _state.runner, _state.verbose, _state.runner.running_trial, _state.log_run_id) @contextlib.contextmanager @@ -287,20 +369,28 @@ def use_run_context(ctx: Optional["_RunContext"]): `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); 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. + 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): @@ -1180,14 +1270,19 @@ def _restore_tune_state(): if verbose: logger.info(f"trial {num_trials} config: {trial_to_run.config}") result = None - # Pin this evaluation call's runner/trial in _propagated_context - # so a worker thread evaluation_function spawns on its own can - # call tune.report() and land on the right trial with no - # trainable-side change (#996 follow-up, second review point 1). - # _install_thread_context_propagation() is what makes a plain - # threading.Thread started inside this call see it. + # 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) + _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): diff --git a/test/tune/test_concurrent_run.py b/test/tune/test_concurrent_run.py index fe8257e5df..e2c2c1703b 100644 --- a/test/tune/test_concurrent_run.py +++ b/test/tune/test_concurrent_run.py @@ -20,6 +20,7 @@ and B's stop_trial() call raises AttributeError on that stale/None runner. """ +import concurrent.futures import threading from unittest import mock @@ -419,3 +420,188 @@ def test_tune_run_spark_setup_failure_restores_state(): ) 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_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}" + ) From 04463eb38f8f7197bdf62c64a70cc32fbac71ab9 Mon Sep 17 00:00:00 2001 From: Amir Fathi Date: Sun, 27 Sep 2026 14:47:18 +0000 Subject: [PATCH 06/10] fix: drop late reports into an already-finished trial (#996 follow-up) thinkall's fourth CHANGES_REQUESTED review, point 2: a _RunContext captured for a trial stays usable after that trial finishes. A worker thread or executor task that captured get_run_context() and reports late, after the trial is already TERMINATED, still wrote through: process_trial_result() overwrote the trial's final metric_analysis and last_result with the stale 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. report() now checks trial.is_finished() before processing and drops the late report instead. New regression test captures a context, lets the trial finish, then reports through the stale context: fails with an uncaught StopIteration on the prior commit, passes now, and asserts last_result/metric_analysis are byte-for-byte unchanged by the late write. Point 1 (a persistent callback/queue worker started before tune.run() never gets context for later items) is not fixed here: replied on the PR with why an automatic fallback is not safe to ship. --- flaml/tune/tune.py | 9 +++++++ test/tune/test_concurrent_run.py | 42 ++++++++++++++++++++++++++++++++ 2 files changed, 51 insertions(+) diff --git a/flaml/tune/tune.py b/flaml/tune/tune.py index 56b5c73430..449685670d 100644 --- a/flaml/tune/tune.py +++ b/flaml/tune/tune.py @@ -574,6 +574,15 @@ def compute_with_config(config): trial = running_trial if running_trial is not None else getattr(runner, "running_trial", None) if not trial: 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. + return None result["training_iteration"] = _next_training_iteration(trial) result["config"] = trial.config if INCUMBENT_RESULT in result["config"]: diff --git a/test/tune/test_concurrent_run.py b/test/tune/test_concurrent_run.py index e2c2c1703b..78886f4510 100644 --- a/test/tune/test_concurrent_run.py +++ b/test/tune/test_concurrent_run.py @@ -569,6 +569,48 @@ def worker(): ) +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 From 936c94759dc7544f7a0fb2c76bacd5cf1331ceeb Mon Sep 17 00:00:00 2001 From: Amir Fathi Date: Sun, 27 Sep 2026 21:33:41 +0000 Subject: [PATCH 07/10] fix: synchronize report admission with trial lifecycle (#996 follow-up) thinkall's fifth CHANGES_REQUESTED review, point 2: trial.is_finished() and runner.process_trial_result() were not atomic. A second, truly concurrent report for the same trial (a trainable's own worker threads reporting for the trial they share, for instance) could pass the is_finished() check while the trial was still running, and if the FIRST report's scheduler decision finished the trial before the second one reached process_trial_result(), the second one still wrote through, silently replacing the trial's real final result with a stale one. Added a per-trial admission lock and moved the actual admission decision (the check process_trial_result() acts on) inside it, so the re-check and the write are now one atomic step. The existing fast-path check right after resolving the trial is unchanged; it just stops being the only line of defense. New test forces the exact interleaving deterministically (events, not a sleep): worker B is paused, via a patched _next_training_iteration, right after its own is_finished() check returns False, until worker A's entire report() call, including the scheduler decision that terminates the trial, has completed. Verified this reproduces the bug on the prior commit (B's stale report wins, 2.0 instead of A's 1.0) and is fixed here. Point 1 (a persistent queue worker started before tune.run() lacking per-dispatch context) is restated from the fourth review. Automatic propagation for it is still not being added, for the reason already given: nothing at report()-call time can reconstruct which trial a delayed item was produced for. The compatibility route the review also names already exists (get_run_context()/use_run_context()); added a test exercising it on exactly the shape described, a worker thread started before tune.run() and reused across two trials. --- flaml/tune/tune.py | 59 ++++++++++- test/tune/test_concurrent_run.py | 176 +++++++++++++++++++++++++++++++ 2 files changed, 230 insertions(+), 5 deletions(-) diff --git a/flaml/tune/tune.py b/flaml/tune/tune.py index 449685670d..f58930574f 100644 --- a/flaml/tune/tune.py +++ b/flaml/tune/tune.py @@ -304,6 +304,42 @@ def _flaml_tune_wrapped(*a, **kw): _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 + def _next_training_iteration(trial) -> int: """Return trial's next training_iteration, as a counter shared by every @@ -582,6 +618,9 @@ def compute_with_config(config): # 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["training_iteration"] = _next_training_iteration(trial) result["config"] = trial.config @@ -589,11 +628,21 @@ def compute_with_config(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 + 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 + runner.process_trial_result(trial, result) + if verbose > 2: + logger.info(f"result: {result}") + if trial.is_finished(): + raise StopIteration def run( diff --git a/test/tune/test_concurrent_run.py b/test/tune/test_concurrent_run.py index 78886f4510..adbafe1a4e 100644 --- a/test/tune/test_concurrent_run.py +++ b/test/tune/test_concurrent_run.py @@ -21,6 +21,7 @@ """ import concurrent.futures +import queue import threading from unittest import mock @@ -647,3 +648,178 @@ def eval_logging_executor(config): "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 (via a + patched _next_training_iteration, the call immediately following it) + 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(). + """ + + 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_next_iteration = tune.tune._next_training_iteration + + def gated_next_iteration(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 orig_next_iteration(trial_obj) + + 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._next_training_iteration", side_effect=gated_next_iteration): + 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}" + ) From 0a055a6aa1556bf023d16f7ec16f4f766d01102f Mon Sep 17 00:00:00 2001 From: Amir Fathi Date: Mon, 28 Sep 2026 03:17:19 +0000 Subject: [PATCH 08/10] fix: address sixth CHANGES_REQUESTED review on #996 thread-safety fix thinkall's sixth review, three real gaps fixed: 1. stop_trial() and report()'s admission did not share a lock. A straggling background-thread report could still be mutating trial.last_result/status via process_trial_result() at the same instant run()'s own loop called stop_trial() on the same trial. stop_trial() now takes the same per-trial _admission_lock_for report() already used. 2. training_iteration was allocated before the admission lock was acquired, so two concurrent reports for one trial could be handed iterations in one order and reach process_trial_result() in the other. Allocation moved inside the same critical section as admission. 3. _logger_level_enter() saved logger.getEffectiveLevel() (the resolved level, walking up the logger hierarchy) instead of logger.level (the logger's own, possibly NOTSET, level), so a logger that was inheriting ended up pinned to an explicit level after every run finished. Restores logger.level now. The review's fourth point, automatic context propagation for a generic pre-started queue worker, is not implemented for the reason given in the third and fourth review rounds: report() cannot tell which trial a delayed queue item belongs to once the runner has moved on to a later trial, and a fallback that guesses would silently attribute a report to the wrong trial, worse than today's silent drop. The documented get_run_context()/use_run_context() API already covers this case explicitly; added a test for the specific pre-started, reused-across-trials shape the review describes. Five new regression tests, each with a negative control confirming it fails on the pre-fix code and passes on this commit. Signed-off-by: Amir Fathi --- flaml/tune/trial_runner.py | 43 +++-- flaml/tune/tune.py | 34 +++- test/tune/test_concurrent_run.py | 273 +++++++++++++++++++++++++++++-- 3 files changed, 327 insertions(+), 23 deletions(-) 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 f58930574f..8558eae8af 100644 --- a/flaml/tune/tune.py +++ b/flaml/tune/tune.py @@ -116,10 +116,28 @@ def filter(self, record): 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.getEffectiveLevel() + _logger_pristine_level = logger.level _active_log_levels.append(level) logger.setLevel(min(_active_log_levels)) @@ -622,7 +640,6 @@ def compute_with_config(config): # 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["training_iteration"] = _next_training_iteration(trial) result["config"] = trial.config if INCUMBENT_RESULT in result["config"]: del result["config"][INCUMBENT_RESULT] @@ -638,6 +655,19 @@ def compute_with_config(config): # 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}") diff --git a/test/tune/test_concurrent_run.py b/test/tune/test_concurrent_run.py index adbafe1a4e..a076362806 100644 --- a/test/tune/test_concurrent_run.py +++ b/test/tune/test_concurrent_run.py @@ -21,6 +21,7 @@ """ import concurrent.futures +import logging import queue import threading from unittest import mock @@ -29,6 +30,7 @@ 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(): @@ -662,13 +664,24 @@ def test_concurrent_reports_for_same_trial_admit_atomically(): stale one. Forced deterministically, with events rather than a sleep: worker B is - paused right after its own is_finished() check returns False (via a - patched _next_training_iteration, the call immediately following it) - 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(). + 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: @@ -695,13 +708,14 @@ def on_trial_remove(self, runner, trial): b_ready = threading.Event() a_finished = threading.Event() - orig_next_iteration = tune.tune._next_training_iteration + orig_admission_lock_for = tune.tune._admission_lock_for - def gated_next_iteration(trial_obj): + 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 orig_next_iteration(trial_obj) + return lock def eval_concurrent_reporters(config): ctx = tune.get_run_context() @@ -731,7 +745,7 @@ def worker_b(): thread_b.join(timeout=5) return None - with mock.patch("flaml.tune.tune._next_training_iteration", side_effect=gated_next_iteration): + 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)}, @@ -823,3 +837,240 @@ def eval_via_prestarted_queue(config): "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) From 865bd8c35ada19e5dff59098bd0d3507a1150364 Mon Sep 17 00:00:00 2001 From: Amir Fathi Date: Mon, 28 Sep 2026 06:59:57 +0000 Subject: [PATCH 09/10] fix: warn once when a context-less report is dropped (#996 seventh review) A generic queue/callback worker started before tune.run() has neither a propagated context nor a runner, so its tune.report() calls silently returned without recording the metric. Automatic attribution is still not added (SequentialTrialRunner.step() reassigns running_trial every step, so a live-resolved fallback would attribute a delayed report to whichever trial happens to be running by the time it is handled, not the one it was produced for) but the drop is no longer silent: it logs once, and the record bypasses the run-scoped filter so it reaches the active run's own log file instead of being swallowed by it. New test starts a legacy queue worker before tune.run(), an unchanged caller with no get_run_context()/use_run_context(), and asserts the trial's result stays untouched (None) while the warning appears exactly once in the run's log file. Fails on the prior commit with an AttributeError (no _context_less_report_warned to patch), passes on this one. Verified in Docker (python:3.11-slim): test/tune/test_concurrent_run.py 18/18 pass; test/tune/ minus test_tune.py (needs xgboost, pre-existing gap) 98 passed, 1 skipped; black 23.3.0 and ruff 0.0.261 clean. Signed-off-by: Amir Fathi --- flaml/tune/tune.py | 98 +++++++++++++++++++++++++++++++- test/tune/test_concurrent_run.py | 79 +++++++++++++++++++++++++ 2 files changed, 175 insertions(+), 2 deletions(-) diff --git a/flaml/tune/tune.py b/flaml/tune/tune.py index 8558eae8af..2b31838d08 100644 --- a/flaml/tune/tune.py +++ b/flaml/tune/tune.py @@ -93,6 +93,20 @@ class _RunScopedFilter(logging.Filter): 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. + + A record carrying `record.flaml_tune_unscoped = True` bypasses the + match and reaches every active run's handler regardless of which + thread emitted it (#996 follow-up, seventh review). This exists for + exactly one caller, the context-less-report warning below: a record + warning that FLAML could not attribute a report to any run has no run + of its own to match against, and Python's logging module only falls + back to printing a record nobody's handler wants (`logging.lastResort`) + when a logger has NO handler at all, not when every handler's filter + happens to reject it, so without this the warning would be silently + swallowed the same way the report it describes is, any time a run is + active with verbose > 0. Verified directly: a filtered handler present + on the logger suppresses lastResort with zero output anywhere, even + though the handler never actually emitted the record. """ def __init__(self, run_id): @@ -100,7 +114,7 @@ def __init__(self, run_id): self._run_id = run_id def filter(self, record): - return _state.log_run_id is self._run_id + return getattr(record, "flaml_tune_unscoped", False) or _state.log_run_id is self._run_id # Bookkeeping for the shared logger's OWN level (logger.setLevel), which is a @@ -547,6 +561,43 @@ def best_iteration(self) -> List[str]: return None +# Guards the one-time warning below (#996 follow-up, seventh 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 every report() it +# makes 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 the +# warning fires once per process rather than once per call: the message is +# the same regardless of which call produced it, and a caller stuck in +# that state needs to see it once, not on every iteration. +_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 + 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). The report is dropped, not recorded against any " + "trial, and this will not be logged again. A worker thread started before " + "tune.run() is called, such as a persistent queue consumer, cannot be " + "attributed to a trial automatically, since resolving it live would risk " + "attributing the report to whichever trial happens to be running by the " + "time it is handled, not the one it was produced for. Capture " + "get_run_context() in the code that queues the work and attach it with " + "use_run_context() where the item is processed.", + extra={"flaml_tune_unscoped": True}, + ) + + def report(_metric=None, **kwargs): """A function called by the HPO application to report final or intermediate results. @@ -591,14 +642,42 @@ def compute_with_config(config): _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, is a case this cannot cover automatically and + is unsupported by design, not an oversight: such a thread has no + _propagated_context (that value did not exist yet when the thread + started) and made no get_run_context()/use_run_context() call of its + own, so there is nothing here that says which trial its report + belongs to. Guessing from the runner's current trial would attribute + the report to whichever trial happens to be running by the time it is + handled, which is very likely not the trial the report was actually + for (#996 follow-up, fourth, fifth and sixth review rounds). The + report is dropped instead, and this logs a one-time warning pointing + at get_run_context()/use_run_context() as the fix. Callers who need + this to work should capture get_run_context() where the work is + produced and attach it with use_run_context() where it is processed. """ 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 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 @@ -616,7 +695,13 @@ def compute_with_config(config): return session.report(metrics={"metric": _metric, **kwargs}) except ImportError: - # calling tune.report() outside tune.run() + # 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: @@ -627,6 +712,15 @@ def compute_with_config(config): # 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: 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 diff --git a/test/tune/test_concurrent_run.py b/test/tune/test_concurrent_run.py index a076362806..ba4ad7a391 100644 --- a/test/tune/test_concurrent_run.py +++ b/test/tune/test_concurrent_run.py @@ -28,6 +28,7 @@ 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 @@ -1074,3 +1075,81 @@ def test_logger_level_restored_as_inherited_not_explicit(): ) finally: logger.setLevel(saved_level) + + +def test_context_less_report_from_unchanged_legacy_queue_worker_warns_without_corrupting(tmp_path): + """Follow-up to #996, seventh review: 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), still has its tune.report() silently dropped: automatic + attribution is not being added, for the same reason given in the + fourth, fifth and sixth review rounds (SequentialTrialRunner.step() + reassigns running_trial every step, so a live-resolved fallback would + attribute a delayed report to whichever trial happens to be running by + the time it is handled, not the one it was produced for). + + What changed this round: the drop is no longer silent. It logs once, + and the message survives even while a verbose run is active with its + own log-file handler attached, which is the exact case + _RunScopedFilter would otherwise swallow it in (a handler whose filter + rejects a record still counts toward Python's `found` handler count, + so `logging.lastResort` never fires either; verified directly while + building this fix). + """ + 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 (fail to) report it + 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 + with mock.patch.object(tune_module, "_context_less_report_warned", False): + analysis = tune.run( + eval_via_legacy_queue, + config={"x": tune.uniform(0, 1)}, + points_to_evaluate=[{"x": 5.0}], + metric="metric", + mode="min", + num_samples=1, + verbose=1, + log_file_name=log_path, + ) + finally: + # Put `stop` and join unconditionally, including if the mock.patch + # setup itself raised (as it does against a tune.py that predates + # this fix, which has no `_context_less_report_warned` attribute to + # patch): 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) + + assert analysis.trials[0].last_result is None, ( + "the context-less report should have been dropped, leaving the trial's result " + f"untouched (None), not corrupted with a value; got {analysis.trials[0].last_result!r}" + ) + + log_text = open(log_path).read() + assert ( + "no active run to attribute it to" in log_text + ), f"expected the context-less-report warning in the run's own log file, got: {log_text!r}" + assert ( + log_text.count("no active run to attribute it to") == 1 + ), "the warning should fire once per process, not once per dropped report" From 2876ab53067255604574c7c174f586ac65d4c5ff Mon Sep 17 00:00:00 2001 From: Amir Fathi Date: Wed, 30 Sep 2026 10:32:07 +0000 Subject: [PATCH 10/10] fix: preserve legacy worker reporting and isolate diagnostics (#996 eighth review) --- flaml/tune/tune.py | 403 +++++++++++++++++++++---------- test/tune/test_concurrent_run.py | 221 +++++++++++++---- 2 files changed, 459 insertions(+), 165 deletions(-) diff --git a/flaml/tune/tune.py b/flaml/tune/tune.py index 2b31838d08..8ab972f34c 100644 --- a/flaml/tune/tune.py +++ b/flaml/tune/tune.py @@ -94,19 +94,17 @@ class _RunScopedFilter(logging.Filter): since `_state` is thread-local and Python's logging dispatch runs synchronously on the emitting thread. - A record carrying `record.flaml_tune_unscoped = True` bypasses the - match and reaches every active run's handler regardless of which - thread emitted it (#996 follow-up, seventh review). This exists for - exactly one caller, the context-less-report warning below: a record - warning that FLAML could not attribute a report to any run has no run - of its own to match against, and Python's logging module only falls - back to printing a record nobody's handler wants (`logging.lastResort`) - when a logger has NO handler at all, not when every handler's filter - happens to reject it, so without this the warning would be silently - swallowed the same way the report it describes is, any time a run is - active with verbose > 0. Verified directly: a filtered handler present - on the logger suppresses lastResort with zero output anywhere, even - though the handler never actually emitted the record. + 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): @@ -114,7 +112,7 @@ def __init__(self, run_id): self._run_id = run_id def filter(self, record): - return getattr(record, "flaml_tune_unscoped", False) or _state.log_run_id is self._run_id + return _state.log_run_id is self._run_id # Bookkeeping for the shared logger's OWN level (logger.setLevel), which is a @@ -373,6 +371,68 @@ def _admission_lock_for(trial) -> threading.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. @@ -561,16 +621,40 @@ def best_iteration(self) -> List[str]: return None -# Guards the one-time warning below (#996 follow-up, seventh review). A +# 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 every report() it -# makes 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 the -# warning fires once per process rather than once per call: the message is -# the same regardless of which call produced it, and a caller stuck in -# that state needs to see it once, not on every iteration. +# 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() @@ -583,18 +667,18 @@ def _warn_context_less_report_once() -> None: if _context_less_report_warned: return _context_less_report_warned = True - logger.warning( + _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). The report is dropped, not recorded against any " - "trial, and this will not be logged again. A worker thread started before " - "tune.run() is called, such as a persistent queue consumer, cannot be " - "attributed to a trial automatically, since resolving it live would risk " - "attributing the report to whichever trial happens to be running by the " - "time it is handled, not the one it was produced for. Capture " - "get_run_context() in the code that queues the work and attach it with " - "use_run_context() where the item is processed.", - extra={"flaml_tune_unscoped": True}, + "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.", ) @@ -644,19 +728,28 @@ def compute_with_config(config): (#996 follow-up, second review point 1). See _RunContext. A thread started before tune.run() is called, such as a persistent - queue/callback worker, is a case this cannot cover automatically and - is unsupported by design, not an oversight: such a thread has no - _propagated_context (that value did not exist yet when the thread - started) and made no get_run_context()/use_run_context() call of its - own, so there is nothing here that says which trial its report - belongs to. Guessing from the runner's current trial would attribute - the report to whichever trial happens to be running by the time it is - handled, which is very likely not the trial the report was actually - for (#996 follow-up, fourth, fifth and sixth review rounds). The - report is dropped instead, and this logs a one-time warning pointing - at get_run_context()/use_run_context() as the fix. Callers who need - this to work should capture get_run_context() where the work is - produced and attach it with use_run_context() where it is processed. + 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 @@ -674,6 +767,11 @@ def compute_with_config(config): # 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: @@ -682,91 +780,121 @@ def compute_with_config(config): runner = ctx.runner verbose = ctx.verbose running_trial = ctx.running_trial - if use_ray: - try: - from ray import __version__ as ray_version + 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 + if ray_version.startswith("1."): + from ray import tune + + return tune.report(_metric, **kwargs) + else: # ray>=2 + from ray.air import session - 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. + 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 - 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: 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 + # 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( @@ -1024,6 +1152,15 @@ def easy_objective(config): _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. @@ -1040,8 +1177,20 @@ def _restore_tune_state(): _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) @@ -1436,6 +1585,14 @@ def _restore_tune_state(): 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 diff --git a/test/tune/test_concurrent_run.py b/test/tune/test_concurrent_run.py index ba4ad7a391..db8a13f63e 100644 --- a/test/tune/test_concurrent_run.py +++ b/test/tune/test_concurrent_run.py @@ -1077,24 +1077,32 @@ def test_logger_level_restored_as_inherited_not_explicit(): logger.setLevel(saved_level) -def test_context_less_report_from_unchanged_legacy_queue_worker_warns_without_corrupting(tmp_path): - """Follow-up to #996, seventh review: a generic queue/callback worker - started before any tune.run() call exists, that never calls +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), still has its tune.report() silently dropped: automatic - attribution is not being added, for the same reason given in the - fourth, fifth and sixth review rounds (SequentialTrialRunner.step() - reassigns running_trial every step, so a live-resolved fallback would - attribute a delayed report to whichever trial happens to be running by - the time it is handled, not the one it was produced for). - - What changed this round: the drop is no longer silent. It logs once, - and the message survives even while a verbose run is active with its - own log-file handler attached, which is the exact case - _RunScopedFilter would otherwise swallow it in (a handler whose filter - rejects a record still counts toward Python's `found` handler count, - so `logging.lastResort` never fires either; verified directly while - building this fix). + 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() @@ -1113,43 +1121,172 @@ def worker(): def eval_via_legacy_queue(config): work_queue.put(config["x"]) - work_queue.join() # wait for the pre-started worker to (fail to) report it + 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 - with mock.patch.object(tune_module, "_context_less_report_warned", False): - analysis = tune.run( - eval_via_legacy_queue, - config={"x": tune.uniform(0, 1)}, - points_to_evaluate=[{"x": 5.0}], - metric="metric", - mode="min", - num_samples=1, - verbose=1, - log_file_name=log_path, - ) + 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, including if the mock.patch - # setup itself raised (as it does against a tune.py that predates - # this fix, which has no `_context_less_report_warned` attribute to - # patch): 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. + # 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) - assert analysis.trials[0].last_result is None, ( - "the context-less report should have been dropped, leaving the trial's result " - f"untouched (None), not corrupted with a value; got {analysis.trials[0].last_result!r}" + 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 log_text - ), f"expected the context-less-report warning in the run's own log file, got: {log_text!r}" + "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 ( - log_text.count("no active run to attribute it to") == 1 - ), "the warning should fire once per process, not once per dropped report" + "MARKER_DIAGNOSTICS_ONLY" not in log_text + ), f"a _diagnostics_logger record leaked into a run-owned log file: {log_text!r}"