Skip to content

Commit 04463eb

Browse files
committed
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.
1 parent 16773f5 commit 04463eb

2 files changed

Lines changed: 51 additions & 0 deletions

File tree

‎flaml/tune/tune.py‎

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -574,6 +574,15 @@ def compute_with_config(config):
574574
trial = running_trial if running_trial is not None else getattr(runner, "running_trial", None)
575575
if not trial:
576576
return None
577+
if trial.is_finished():
578+
# A late report from a background thread or executor task whose
579+
# captured _RunContext outlived its trial (#996 follow-up, fourth
580+
# review point 2): the trial's final result is already recorded,
581+
# and process_trial_result() would overwrite it with this stale
582+
# value, plus the is_finished() check below would then raise
583+
# StopIteration into a caller that never expected it (unlike the
584+
# trainable's own control-flow loop, which does). Drop it instead.
585+
return None
577586
result["training_iteration"] = _next_training_iteration(trial)
578587
result["config"] = trial.config
579588
if INCUMBENT_RESULT in result["config"]:

‎test/tune/test_concurrent_run.py‎

Lines changed: 42 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -569,6 +569,48 @@ def worker():
569569
)
570570

571571

572+
def test_late_report_via_stale_context_does_not_corrupt_finished_trial():
573+
"""Follow-up to #996, fourth review point 2: a _RunContext captured for
574+
a trial stays usable after that trial finishes. A worker thread that
575+
captured get_run_context() and reports late, after tune.run() has
576+
already returned and the trial is TERMINATED, used to still write
577+
through: process_trial_result() overwrote the trial's already-final
578+
metric_analysis/last_result with the late value, and report()'s own
579+
trailing `if trial.is_finished(): raise StopIteration` (the normal
580+
scheduler-stop signal for the CURRENT report) then raised into the late
581+
caller too, which has no reason to expect it the way a trainable's own
582+
control-flow loop does.
583+
"""
584+
captured = {}
585+
586+
def eval_capturing(config):
587+
captured["ctx"] = tune.get_run_context()
588+
return {"metric": 1.0}
589+
590+
analysis = tune.run(
591+
eval_capturing,
592+
config={"x": tune.uniform(0, 1)},
593+
metric="metric",
594+
mode="min",
595+
num_samples=1,
596+
verbose=0,
597+
)
598+
trial = analysis.trials[0]
599+
assert trial.is_finished(), "expected the trial to be TERMINATED once tune.run() returns"
600+
last_result_before = dict(trial.last_result)
601+
metric_analysis_before = {k: dict(v) for k, v in trial.metric_analysis.items()}
602+
603+
with tune.use_run_context(captured["ctx"]):
604+
tune.report(metric=999.0)
605+
606+
assert (
607+
trial.last_result == last_result_before
608+
), f"a late report corrupted the finished trial's last_result: {trial.last_result}"
609+
assert (
610+
trial.metric_analysis == metric_analysis_before
611+
), f"a late report corrupted the finished trial's metric_analysis: {trial.metric_analysis}"
612+
613+
572614
def test_tune_log_records_from_executor_worker_reach_run_log(tmp_path):
573615
"""Same as above (third review point 4), for a task submitted to a
574616
ThreadPoolExecutor rather than a plain Thread, since the two are

0 commit comments

Comments
 (0)