[train] Add optional step and weight-sync timeouts that raise typed errors - #2381
dyurk-lila wants to merge 1 commit into
Conversation
Bound awaited sections, Ray waits, and inference control-plane requests. Cover the example async trainer, preserve the primary error during bounded cleanup, and keep the deadline tests focused on real call paths and distinct failure modes.
There was a problem hiding this comment.
Code Review
This pull request introduces driver-side step and weight sync deadlines to prevent training steps from hanging indefinitely. It adds step_timeout_s and weight_sync_timeout_s configuration options to TrainerConfig, which raise StepTimeoutError or WeightSyncTimeoutError when exceeded. A new deadline utility module is introduced to manage these budgets, replacing standard ray.get calls with bounded deadline.ray_get calls across workers and dispatchers. Additionally, control-plane requests to inference servers are now bounded by SKYRL_INFERENCE_CONTROL_PLANE_TIMEOUT_S, raising InferenceServerTimeoutError on timeout. Relevant documentation, configuration validation, and comprehensive unit tests have been added to support these changes. I have no feedback to provide as there are no review comments.
|
| await self._vllm_metrics_scraper.start("vllm/train") | ||
| self._vllm_metrics_scraper.pause() | ||
| with Timer("step", self.all_timings): | ||
| async with self._step_deadline(), Timer("step", self.all_timings): |
There was a problem hiding this comment.
Unbounded waits remain inside steps
With policy profiling enabled, a step calls _profiler_step(), whose worker dispatch still uses an unbounded synchronous ray.get. If that worker hangs, it blocks the event loop, so the new step deadline cannot raise StepTimeoutError and the driver remains stuck. Checkpoint cleanup has the same gap through run_on_each_node().
Knowledge Base Used: Trainer execution and evaluation
| async with self._step_deadline(): | ||
| with self._phase_gauge.timed_phase("save_checkpoints", self.all_timings): | ||
| await asyncio.to_thread(self.save_checkpoints) |
There was a problem hiding this comment.
If the deadline expires during asyncio.to_thread(self.save_checkpoints), it cancels the await but not the save thread. That thread can keep writing after StepTimeoutError propagates, and the pending-save drain does not wait for it. The base save writes the latest-step marker before the fully asynchronous save writes fully_async_state.pt, which resume requires, so a timed-out save can leave a checkpoint that cannot be resumed. The HF save has the same continuing-thread problem.
Knowledge Base Used: Trainer execution and evaluation
| remaining_s = deadline.remaining() | ||
| if remaining_s == 0: | ||
| logger.warning(f"Skipping {description} during cleanup: the step deadline has expired") | ||
| raise |
There was a problem hiding this comment.
Expired sync skips engine recovery
If a weight-sync deadline expires after an external inference server has been paused or put to sleep, this branch skips the matching resume_generation or wake_up call. The trainer then propagates the error without restoring that external server, leaving it paused or asleep after the training job exits.
Knowledge Base Used:
| if value <= 0: | ||
| raise ValueError(f"trainer.{name} must be > 0, got {value}") |
There was a problem hiding this comment.
There was a problem hiding this comment.
Cursor Bugbot has reviewed your changes using default effort and found 2 potential issues.
❌ Bugbot Autofix is OFF. To automatically fix reported issues with cloud agents, have a team admin enable autofix in the Cursor dashboard.
Reviewed by Cursor Bugbot for commit 30669e0. Configure here.
|
|
||
| pbar.close() | ||
| if self.cfg.trainer.ckpt_interval > 0: | ||
| with Timer("save_checkpoints", self.all_timings): |
There was a problem hiding this comment.
Periodic saves omit timeout budget
Medium Severity
Interval checkpoint and HF saves in the example async trainer sit outside the step scope and never enter _step_deadline(), unlike the final saves in the same loop. With step_timeout_s set, a hung save_checkpoints or save_models call still blocks the driver indefinitely.
Reviewed by Cursor Bugbot for commit 30669e0. Configure here.
| trained_steps_this_epoch = self.async_train_dataloader.num_trained() // self.mini_batch_size | ||
| for _step_idx in range(self.global_step, (1 + epoch) * self.num_steps_per_epoch + 1): | ||
| with Timer("step", self.all_timings): | ||
| async with self._step_deadline(), Timer("step", self.all_timings): |
There was a problem hiding this comment.
Early-exit saves reuse leftover budget
Low Severity
The sample_full_batch epoch-exhaustion checkpoint and HF saves still run inside the step’s _step_deadline(), so they inherit whatever time is left after collection. Other fully-async saves open a fresh budget. A long collect followed by this early-exit path can fail the save with StepTimeoutError even though the step’s work already finished.
Reviewed by Cursor Bugbot for commit 30669e0. Configure here.


What does this PR do?
Adds two optional driver-side budgets,
trainer.step_timeout_sandtrainer.weight_sync_timeout_s. A hung training step becomes a typed, picklableStepTimeoutError/WeightSyncTimeoutErrorthat names the step, stage and operation.TLDR: the deadline enters
asyncio.timeout, so every await in the step (generation, inference-enginepause/resumeandsleep/wake_up, the fully-async buffer wait,to_threadsaves) is cancelled on expiry. Theray.gets behind a step'sWorkerDispatchcalls andPPORayActorGroupoffload/backload go throughdeadline.ray_get. Both trainer keys default toNone; control-plane HTTP requests have a separate 600 s default timeout.Usage
step_timeout_sbudget. Saves inside a sync-trainer step share the step's budget.SKYRL_INFERENCE_CONTROL_PLANE_TIMEOUT_Sbounds each inference-server control request (default 600 s;<= 0disables it). A request timeout raises picklableInferenceServerTimeoutError. It does not apply to data-plane generation requests.How it works
skyrl/train/utils/deadline.py:step_deadline(global_step, budget_s, error_cls)sets a_Deadlinein aContextVarand wraps its body inasyncio.timeout(budget_s). It is a no-op forNone. A nested deadline can only shorten its parent: if the parent expires first, the parent, and its error class, stay in effect.Timers anddeadline.operation(...)scopes. Each records itself on the deadline, innermost first, and the deadline raiseserror_cls(step, stage, operation or "await", …). ATimeoutErrorfrom the body that isn't this deadline's passes through unchanged.ray_get(refs, operation)callsray.get(refs, timeout=remaining())and turnsGetTimeoutErrorinto the deadline's error.stage(name)names the current stage, andoperation(name)the call awaited within it:RemoteInferenceClient._call_all_serversnames its endpoint and the sync trainer namesgenerate, so a timeout during a weight sync reads e.g.stage='sync_weights', operation='/pause'.Timer.__enter__/__aenter__enters it with the timer's message, so every existingTimer("…")/timed_phase("…")names its stage with no new bookkeeping.asyncio.to_threadcopies the context, soray_getin worker threads sees both the deadline and the stage.StepTimeoutError(RuntimeError)passes all five fields tosuper().__init__, so default pickling round-trips them. It is deliberately not aTimeoutError, soexcept TimeoutErrorretry paths (e.g.RemoteInferenceClient's) can't swallow it.WeightSyncTimeoutErrorsubclasses it.cleanup_preserving_primary's failure-path cleanup (theresume_generation/wake_upinsave_weights_for_sampler) runs underasyncio.timeout(deadline.remaining()). It is skipped, with a warning, once the deadline has expired, and a cleanup that hangs past it adds a note to the primary exception. Without a deadline it is unbounded, as before.asyncio.timeoutcancels only once, so without this a wedged engine would hang the unwind after the deadline fired.Trainer scopes: the step block (
async with self._step_deadline(), Timer("step", …)), each weight sync, including the initial one (_weight_sync_deadline()), and each save outside a step. The synchronous, fully-async and example async trainers all enter these scopes.RemoteInferenceClientuses anaiohttp.ClientTimeoutfor each control-plane request, including LoRA load/unload. This bounds calls made outside a training step too.validate_cfg:<= 0values andweight_sync_timeout_s > step_timeout_s;SKYRL_WORKER_NCCL_TIMEOUT_IN_S, because the driver would then give up before the NCCL watchdog reports which collective is stuck.The keys' docstrings render in the config API reference. The troubleshooting page gets a "Step timeouts" section.
Alternatives considered:
timeout=arguments: they would have to be threaded through every dispatch signature, and they can't express "whatever is left of this step".os._exitpattern the previous PR removes.Test plan
tests/train/test_deadline.py:ray_getuses the remaining budget and the enclosingTimerstage.backload_to_gputimes out through the actor-group call site; both error classes survive pickling."await"fallback; an inner weight-sync deadline raisesWeightSyncTimeoutError.TimeoutErrors pass through. After expiry, cleanup is skipped; a cleanup that hangs after a body error is cut off and adds a note to the original error.tests/train/test_config.py: parameterized rejection of invalid values and the NCCL warning.test_remote_inference_client.py: a real mock-server control request that sleeps past a short budget raises a picklableInferenceServerTimeoutError.Validation (whole stack)
Current CPU revision: PR61 trainer tests: 19 passed. PR63 focused stack tests: 322 passed. Ruff passes on all changed Python files. A full-suite rerun stopped around the existing Ray dispatch tests; the standalone dispatch test also timed out after 120 s. A prior stack revision passed the full CPU command (
RAY_ADDRESS=local uv run --isolated --extra skyrl-train --extra dev pytest -q tests/train/ tests/backends/skyrl_train/ --ignore=tests/backends/skyrl_train/gpu -m "not vllm"): 2020 passed, 32 skipped, 6 deselected.Prior GPU end-to-end validation (earlier stack revision; 1× H100 node, 4 GPUs, Qwen2.5-0.5B-Instruct on GSM8K, FSDP + vLLM 0.30). A small driver runs the trainer as a Ray task and inspects the exception
ray.getraises on the driver. Fault injection wrapsgenerator.generate.step_timeout_s=1800,weight_sync_timeout_s=900, ckpt at step 2generateraises after 40 callsRayTaskError(RuntimeError)with the worker traceback, notes['raised in background generation worker']step_timeout_s=2StepTimeoutError(1, 'wait_for_generation_buffer', 'await', 2, 2.0008)generateraises on call 2RayTaskError(RuntimeError), notes['raised in background generator']weight_sync_timeout_s=0.01WeightSyncTimeoutError(0, 'sync_weights', 'backload_to_gpu', 0.01, …)step_timeout_s=3StepTimeoutError(1, 'generate', 'generate', 3, 3.0005)Before this stack, the injected-failure rows end in
os._exit(1)/sys.exit(1), so the driver sees the worker process die.Prior GPU CI subset (
-m "not (integrations or megatron or mooncake)"), run on an earlier stack revision and on unmodifiedmainin the same environment:maintest_save_weights_for_sampler.pytest_offload_kv_weight_sync.pyinference_servers/test_new_inference_generation.pytest_worker_dispatch_offload.pytest_trainer_full_checkpointing.py -k fsdpinference_servers/test_weight_sync.pytest_weight_sync.pyfails the same 4 tests onmain, for the same two reasons on both sides:AttributeError: 'types.SimpleNamespace' object has no attribute 'speculative_config', and a vLLM engine-core init failure.test_update_weights_rdthits one reason on each side.🤖 Generated with Claude Code
Note
Medium Risk
Touches core training loops, Ray dispatch waits, and inference control-plane HTTP; mis-tuned timeouts could abort long but healthy steps, though defaults remain off (None).
Overview
Adds optional driver-side time budgets so hung training no longer blocks forever. Set
trainer.step_timeout_s(whole step: generation through weight sync) and optionallytrainer.weight_sync_timeout_s(must be ≤ step budget); when exceeded the driver raises picklableStepTimeoutError/WeightSyncTimeoutErrorwith step, stage (Timername), and operation.New
skyrl.train.utils.deadlinewraps steps inasyncio.timeout, boundsray.getviadeadline.ray_getacrossWorkerDispatchand worker offload/backload, and tags awaits withdeadline.operation. Sync, fully-async, and example async trainers enter_step_deadline()/_weight_sync_deadline()around steps, weight syncs, and out-of-step checkpoint/HF saves.Timerregisters stages for error messages; failure-pathcleanup_preserving_primaryskips or caps cleanup once the deadline expires.Inference control plane calls (pause/resume, weight updates, LoRA, etc.) use
SKYRL_INFERENCE_CONTROL_PLANE_TIMEOUT_S(default 600s) and raiseInferenceServerTimeoutError. Config validation rejects invalid timeout pairs and warns if budgets are below the NCCL watchdog. Troubleshooting docs describe usage and NCCL interaction. Only the driver wait is cancelled—workers keep running until teardown.Reviewed by Cursor Bugbot for commit 30669e0. Bugbot is set up for automated code reviews on this repo. Configure here.