From 09d34504f724a0ec29848a15304275cb635ec122 Mon Sep 17 00:00:00 2001 From: FurtherAI Date: Fri, 17 Jul 2026 00:25:31 +0000 Subject: [PATCH 1/7] Omit MoE routes from trajectory parquet --- src/art/utils/trajectory_logging.py | 13 ++++++++++--- 1 file changed, 10 insertions(+), 3 deletions(-) diff --git a/src/art/utils/trajectory_logging.py b/src/art/utils/trajectory_logging.py index 8a1f1701a..325f576d6 100644 --- a/src/art/utils/trajectory_logging.py +++ b/src/art/utils/trajectory_logging.py @@ -8,6 +8,7 @@ from openai.types.chat.chat_completion import Choice import pydantic +from art.openai import ART_MOE_ROUTING_METADATA_KEY from art.trajectories import History, Trajectory, TrajectoryGroup if TYPE_CHECKING: @@ -57,7 +58,9 @@ def _compact_json(value: object) -> str: def _choice_data(item: object) -> dict[str, Any] | None: if isinstance(item, Choice): - raw = item.model_dump(mode="json", warnings="error") + raw = item.model_dump( + mode="json", exclude={ART_MOE_ROUTING_METADATA_KEY}, warnings="error" + ) elif ( isinstance(item, Mapping) and { @@ -69,7 +72,9 @@ def _choice_data(item: object) -> dict[str, Any] | None: ): raw = dict(item) try: - return Choice.model_validate(raw).model_dump(mode="json", warnings="error") + return Choice.model_validate(raw).model_dump( + mode="json", exclude={ART_MOE_ROUTING_METADATA_KEY}, warnings="error" + ) except pydantic.ValidationError: return None else: @@ -78,7 +83,9 @@ def _choice_data(item: object) -> dict[str, Any] | None: if not isinstance(item, Choices): return None raw = item.model_dump(mode="json", warnings="error") - return Choice.model_validate(raw).model_dump(mode="json", warnings="error") + return Choice.model_validate(raw).model_dump( + mode="json", exclude={ART_MOE_ROUTING_METADATA_KEY}, warnings="error" + ) def _history_data(history: History) -> dict[str, Any]: From d607534166171afd6c47674c48c45ed7cb0fac0b Mon Sep 17 00:00:00 2001 From: FurtherAI Date: Fri, 17 Jul 2026 02:20:57 +0000 Subject: [PATCH 2/7] Prioritize stale backlog in pipeline autotuning --- src/art/pipeline_tuner/autotune.py | 27 ++++++++++----------------- 1 file changed, 10 insertions(+), 17 deletions(-) diff --git a/src/art/pipeline_tuner/autotune.py b/src/art/pipeline_tuner/autotune.py index ab4c2b359..24f3ad081 100644 --- a/src/art/pipeline_tuner/autotune.py +++ b/src/art/pipeline_tuner/autotune.py @@ -304,6 +304,16 @@ def _decide(self, stats: TunerWindowStats) -> TunerDecision: if stats.queue_put_wait_frac >= self.config.queue_put_severe_frac: reason = "completed-group queue backpressure is active" + elif predicted_stale_high: + updated = updated.model_copy( + update={ + "num_rollout_workers": self._move_workers( + updated.num_rollout_workers, -1 + ) + } + ) + action = "decrease_workers" + reason = "predicted stale backlog exceeds the freshness target" elif state in { "inference_under_train_under", "inference_balanced_train_under", @@ -317,23 +327,6 @@ def _decide(self, stats: TunerWindowStats) -> TunerDecision: ) action = "increase_workers" reason = "vLLM pressure is low and trainer is underfed" - elif state in { - "inference_under_train_over", - "inference_balanced_train_over", - }: - if ( - updated.min_batch_size >= updated.max_batch_size - and predicted_stale_high - ): - updated = updated.model_copy( - update={ - "num_rollout_workers": self._move_workers( - updated.num_rollout_workers, -1 - ) - } - ) - action = "decrease_workers" - reason = "trainer saturated with predicted stale backlog" elif state == "inference_over_train_over": reason = "both sides are loaded; no throughput-safe online change" From d5c60fe9b4390f580bfe86d55798f55601c8c6ea Mon Sep 17 00:00:00 2001 From: FurtherAI Date: Fri, 17 Jul 2026 04:16:14 +0000 Subject: [PATCH 3/7] Stabilize stale-backlog autotuning --- src/art/metrics.py | 24 ++++++++++ src/art/model.py | 5 ++ src/art/pipeline_trainer/trainer.py | 73 ++++++++++++++++------------- src/art/pipeline_tuner/autotune.py | 56 +++++++++++++++++++--- src/art/pipeline_tuner/config.py | 10 +++- 5 files changed, 127 insertions(+), 41 deletions(-) diff --git a/src/art/metrics.py b/src/art/metrics.py index 5c339ecfa..391e3499f 100644 --- a/src/art/metrics.py +++ b/src/art/metrics.py @@ -276,6 +276,30 @@ class MetricDefinition(pydantic.BaseModel): kind="ratio", higher_is_better=False, ), + MetricDefinition( + key="queue/actual_stale_fraction", + title="Actual stale dequeue fraction", + description="stale groups divided by all groups dequeued for this train step", + kind="ratio", + higher_is_better=False, + ), + MetricDefinition( + key="queue/put_wait_frac", + title="Completed queue wait fraction", + description=( + "worker queue-put wait divided by queue-put wait plus rollout time" + ), + kind="ratio", + higher_is_better=False, + ), + MetricDefinition( + key="queue/put_wait_s", + title="Completed queue wait", + description="worker-seconds spent waiting to enqueue completed rollout groups", + kind="duration", + unit="seconds", + higher_is_better=False, + ), ) PIPELINE_RL_DASHBOARD_DEFAULT_METRICS = tuple( diff --git a/src/art/model.py b/src/art/model.py index 36b17fb4d..32de0a065 100644 --- a/src/art/model.py +++ b/src/art/model.py @@ -348,6 +348,8 @@ def __getattr__(self, name: str) -> Any: "data/cum/num_gradient_steps", "discarded/cum/stale_groups", "discarded/cum/zero_variance_groups", + "discarded/step/stale_groups", + "discarded/step/zero_variance_groups", "discarded/rate/stale_groups", "discarded/rate/zero_variance_groups", "time/step_wall_s", @@ -360,6 +362,9 @@ def __getattr__(self, name: str) -> Any: "pipeline_settings/max_batch_size", "pipeline_settings/target_groups_per_step", "pipeline_settings/queue_maxsize", + "queue/actual_stale_fraction", + "queue/put_wait_frac", + "queue/put_wait_s", "loss/train", "loss/entropy", "loss/kl_div", diff --git a/src/art/pipeline_trainer/trainer.py b/src/art/pipeline_trainer/trainer.py index f25f32fdc..f9bbde919 100644 --- a/src/art/pipeline_trainer/trainer.py +++ b/src/art/pipeline_trainer/trainer.py @@ -56,9 +56,6 @@ from .types import ConfigT, EvalFn, RolloutFn, ScenarioT, SingleRolloutFn # noqa: F401 PIPELINE_STATE_KEY = "_pipeline_trainer" -_ROLLOUT_WALL_TIME_KEY = "_art_rollout_wall_s" -_ACTOR_IDLE_TIME_KEY = "_art_actor_idle_s" -_QUEUE_WAIT_TIME_KEY = "_art_queue_wait_s" _SCORE_FRESHNESS_TAU_STEPS = 8.0 # Rollout critical batch size from the best current GRPO/RLVR evidence. This is # grounded in reported experiments, not a well-validated universal constant. @@ -291,6 +288,8 @@ def __init__( ) self._scenario_source_exhausted = False self._output_queue: asyncio.Queue[TrajectoryGroup | None] | None = None + self._producer_rollout_timings = (0.0, 0.0, 0.0) + self._reported_producer_rollout_timings = (0.0, 0.0, 0.0) self._eval_queue: asyncio.Queue[int] | None = None self._rollout_worker_controller = RolloutWorkerController( self, self.num_rollout_workers @@ -830,9 +829,9 @@ async def _rollout_worker(self, worker_id: int) -> None: if self.state.done: break queue_wait_s = await self._put_output_group(group) - group.metadata[_ROLLOUT_WALL_TIME_KEY] = rollout_wall_s - group.metadata[_QUEUE_WAIT_TIME_KEY] = queue_wait_s - group.metadata[_ACTOR_IDLE_TIME_KEY] = actor_idle_s + queue_wait_s + self._record_producer_rollout_timings( + rollout_wall_s, actor_idle_s + queue_wait_s, queue_wait_s + ) except asyncio.CancelledError: raise except LocalServingUnavailableError: @@ -881,18 +880,19 @@ async def _training_stage(self) -> None: break step_start = time.monotonic() collect_started = time.monotonic() + zero_variance_before = self.state.discarded_zero_variance_groups batch, discarded, saw_sentinel = await self._collect_batch(current_step) trainer_idle_s = time.monotonic() - collect_started + zero_variance_discarded = ( + self.state.discarded_zero_variance_groups - zero_variance_before + ) + dequeued_groups = len(batch) + discarded + zero_variance_discarded self.state.discarded_stale_groups += discarded if discarded: self._status.note_stale(discarded) if not batch: break - actor_wall_s, actor_idle_s, queue_wait_s = ( - self._consume_batch_rollout_timings(batch) - ) - training_policy_step = current_step expected_step = current_step + 1 should_eval_step = self._should_eval_step(expected_step) @@ -957,6 +957,9 @@ async def _training_stage(self) -> None: await self._run_checkpoint_retention(current_step) step_seconds = time.monotonic() - step_start + actor_wall_s, actor_idle_s, queue_wait_s = ( + self._consume_producer_rollout_timings() + ) self._status.note_training_batch( batch, step=current_step, step_seconds=step_seconds ) @@ -973,6 +976,9 @@ async def _training_stage(self) -> None: "discarded/cum/stale_groups": stale_groups, "discarded/cum/zero_variance_groups": zero_variance_groups, "discarded/step/stale_groups": float(discarded), + "discarded/step/zero_variance_groups": float( + zero_variance_discarded + ), "discarded/rate/stale_groups": stale_groups / max(generated_groups_cum, 1.0), "discarded/rate/zero_variance_groups": zero_variance_groups @@ -980,14 +986,14 @@ async def _training_stage(self) -> None: "time/step_wall_s": step_seconds, "time/step_collect_batch_s": trainer_idle_s, "time/step_trainer_idle_s": trainer_idle_s, + "time/step_rollout_s": actor_wall_s, + "time/step_rollout_idle_s": actor_idle_s, + "queue/put_wait_s": queue_wait_s, + "queue/put_wait_frac": queue_wait_s + / max(queue_wait_s + actor_wall_s, 1e-9), + "queue/actual_stale_fraction": discarded / max(dequeued_groups, 1), } metrics.setdefault("time/step_backend_train_s", train_call_elapsed) - if actor_wall_s > 0: - metrics["time/step_rollout_s"] = actor_wall_s - if actor_idle_s > 0: - metrics["time/step_rollout_idle_s"] = actor_idle_s - if queue_wait_s > 0 and actor_wall_s > 0: - metrics["queue/put_wait_frac"] = queue_wait_s / actor_wall_s metrics.update(result.metrics) attachment_metrics, attachment_owns_vllm_metrics = ( self._collect_attachment_train_step_metrics() @@ -1822,21 +1828,22 @@ async def _put_output_group(self, group: TrajectoryGroup) -> float: continue return time.monotonic() - queue_wait_started - def _consume_batch_rollout_timings( - self, batch: list[TrajectoryGroup] - ) -> tuple[float, float, float]: - rollout_wall_s = 0.0 - actor_idle_s = 0.0 - queue_wait_s = 0.0 - for group in batch: - rollout_wall_s += self._pop_float_metadata(group, _ROLLOUT_WALL_TIME_KEY) - actor_idle_s += self._pop_float_metadata(group, _ACTOR_IDLE_TIME_KEY) - queue_wait_s += self._pop_float_metadata(group, _QUEUE_WAIT_TIME_KEY) - return rollout_wall_s, actor_idle_s, queue_wait_s + def _record_producer_rollout_timings( + self, rollout_wall_s: float, actor_idle_s: float, queue_wait_s: float + ) -> None: + previous = self._producer_rollout_timings + self._producer_rollout_timings = ( + previous[0] + rollout_wall_s, + previous[1] + actor_idle_s, + previous[2] + queue_wait_s, + ) - @staticmethod - def _pop_float_metadata(group: TrajectoryGroup, key: str) -> float: - value = group.metadata.pop(key, 0.0) - if isinstance(value, (int, float)): - return float(value) - return 0.0 + def _consume_producer_rollout_timings(self) -> tuple[float, float, float]: + current = self._producer_rollout_timings + previous = self._reported_producer_rollout_timings + self._reported_producer_rollout_timings = current + return ( + max(0.0, current[0] - previous[0]), + max(0.0, current[1] - previous[1]), + max(0.0, current[2] - previous[2]), + ) diff --git a/src/art/pipeline_tuner/autotune.py b/src/art/pipeline_tuner/autotune.py index 24f3ad081..aa89d0471 100644 --- a/src/art/pipeline_tuner/autotune.py +++ b/src/art/pipeline_tuner/autotune.py @@ -97,6 +97,7 @@ def __init__( self._last_decision_step = self._warmup_end_step self._target_candidate: int | None = None self._target_candidate_count = 0 + self._stale_backlog_active = False self._emitted_recommendations: set[str] = set() def on_metric(self, rec: PipelineMetric) -> TunerDecision | None: @@ -164,6 +165,18 @@ def step_values(name: str) -> list[float]: groups = _required_step_values( by_step, window_steps, "data/step_num_groups_trainable" ) + stale_groups = _required_step_values( + by_step, window_steps, "discarded/step/stale_groups" + ) + zero_variance_groups = _required_step_values( + by_step, window_steps, "discarded/step/zero_variance_groups" + ) + rollout_s = sum( + _required_step_values(by_step, window_steps, "time/step_rollout_s") + ) + queue_put_wait_s = sum( + _required_step_values(by_step, window_steps, "queue/put_wait_s") + ) train_capacity_tokens = _required_step_values( by_step, window_steps, "data/step_packed_train_tokens" ) @@ -215,8 +228,11 @@ def step_values(name: str) -> list[float]: vllm_pressure=_vllm_pressure( vllm_metrics, window_start_s=t0, window_end_s=t1 ), - queue_put_wait_frac=_mean(step_values("queue/put_wait_frac")), + queue_put_wait_frac=queue_put_wait_s + / max(queue_put_wait_s + rollout_s, 1e-9), predicted_stale_frac=_mean(step_values("queue/predicted_stale_fraction")), + actual_stale_frac=sum(stale_groups) + / max(sum(groups) + sum(stale_groups) + sum(zero_variance_groups), 1.0), padding_ratio_mean=padding_ratio_mean, ) @@ -298,13 +314,25 @@ def _decide(self, stats: TunerWindowStats) -> TunerDecision: target_changed = ( updated.target_groups_per_step != previous.target_groups_per_step ) - predicted_stale_high = stats.predicted_stale_frac >= self.config.stale_high_frac + stale_backlog_active = self._update_stale_backlog_state(stats) action = "hold" reason = "inside hysteresis band or already balanced" - if stats.queue_put_wait_frac >= self.config.queue_put_severe_frac: - reason = "completed-group queue backpressure is active" - elif predicted_stale_high: + if stale_backlog_active and updated.min_batch_size < updated.max_batch_size: + updated = updated.model_copy( + update={ + "min_batch_size": min( + updated.max_batch_size, + max( + updated.min_batch_size + 1, + round(updated.min_batch_size * 1.15), + ), + ) + } + ) + action = "raise_min_batch_size" + reason = "stale backlog requires dense batches before reducing workers" + elif stale_backlog_active: updated = updated.model_copy( update={ "num_rollout_workers": self._move_workers( @@ -313,7 +341,9 @@ def _decide(self, stats: TunerWindowStats) -> TunerDecision: } ) action = "decrease_workers" - reason = "predicted stale backlog exceeds the freshness target" + reason = "predicted or actual stale backlog exceeds the freshness target" + elif stats.queue_put_wait_frac >= self.config.queue_put_severe_frac: + reason = "completed-group queue backpressure is active" elif state in { "inference_under_train_under", "inference_balanced_train_under", @@ -330,7 +360,7 @@ def _decide(self, stats: TunerWindowStats) -> TunerDecision: elif state == "inference_over_train_over": reason = "both sides are loaded; no throughput-safe online change" - if not target_changed: + if not target_changed and not stale_backlog_active: min_update = self._min_batch_adjustment(updated, stats, state, action) if min_update is not None: updated, action, reason = min_update @@ -351,6 +381,18 @@ def _decide(self, stats: TunerWindowStats) -> TunerDecision: stats=stats, ) + def _update_stale_backlog_state(self, stats: TunerWindowStats) -> bool: + stale_fractions = (stats.predicted_stale_frac, stats.actual_stale_frac) + if self._stale_backlog_active: + self._stale_backlog_active = any( + fraction > self.config.stale_clear_frac for fraction in stale_fractions + ) + else: + self._stale_backlog_active = any( + fraction >= self.config.stale_high_frac for fraction in stale_fractions + ) + return self._stale_backlog_active + def _min_batch_adjustment( self, settings: PipelineTuneSettings, diff --git a/src/art/pipeline_tuner/config.py b/src/art/pipeline_tuner/config.py index de1874fde..472de2325 100644 --- a/src/art/pipeline_tuner/config.py +++ b/src/art/pipeline_tuner/config.py @@ -55,8 +55,9 @@ class PipelineAutotuneConfig(pydantic.BaseModel): trainer_load_over_score: float = pydantic.Field(default=0.04, ge=0.0) vllm_pressure_over_ratio: float = pydantic.Field(default=0.80, ge=0.0) vllm_pressure_under_ratio: float = pydantic.Field(default=0.50, ge=0.0) - queue_put_severe_frac: float = pydantic.Field(default=0.50, ge=0.0, le=1.0) + queue_put_severe_frac: float = pydantic.Field(default=1.0 / 3.0, ge=0.0, le=1.0) stale_high_frac: float = pydantic.Field(default=0.20, ge=0.0, le=1.0) + stale_clear_frac: float = pydantic.Field(default=0.10, ge=0.0, le=1.0) padding_high_frac: float = pydantic.Field(default=0.25, ge=0.0, le=1.0) trainer_min_batch_lower_score: float = pydantic.Field(default=0.15, ge=0.0) recommendation_min_windows: int = pydantic.Field(default=5, ge=1) @@ -78,6 +79,12 @@ class PipelineAutotuneConfig(pydantic.BaseModel): default=0.35, ge=0.0, le=1.0 ) + @pydantic.model_validator(mode="after") + def validate_stale_hysteresis(self) -> "PipelineAutotuneConfig": + if self.stale_clear_frac > self.stale_high_frac: + raise ValueError("stale_clear_frac must be <= stale_high_frac") + return self + class PipelineTuneSettings(pydantic.BaseModel): num_rollout_workers: int = pydantic.Field(ge=1) @@ -127,6 +134,7 @@ class TunerWindowStats(pydantic.BaseModel): vllm_pressure: float = 0.0 queue_put_wait_frac: float = 0.0 predicted_stale_frac: float = 0.0 + actual_stale_frac: float = 0.0 padding_ratio_mean: float = 0.0 From 2682d1f8a137d3defc64d8b408f3b550d505e3fb Mon Sep 17 00:00:00 2001 From: FurtherAI Date: Fri, 17 Jul 2026 08:20:19 +0000 Subject: [PATCH 4/7] Make pipeline min-batch tuning conservative --- src/art/pipeline_tuner/attachment.py | 10 ++- src/art/pipeline_tuner/autotune.py | 105 +++++++++++++++++++-------- src/art/pipeline_tuner/config.py | 12 ++- 3 files changed, 94 insertions(+), 33 deletions(-) diff --git a/src/art/pipeline_tuner/attachment.py b/src/art/pipeline_tuner/attachment.py index 793e0c010..8a1985009 100644 --- a/src/art/pipeline_tuner/attachment.py +++ b/src/art/pipeline_tuner/attachment.py @@ -2,6 +2,7 @@ import asyncio import inspect +import math import time from typing import Any import warnings @@ -334,12 +335,19 @@ def _settings_with_current_queue( ) -> PipelineTuneSettings: return settings.model_copy( update={ + "min_batch_size": max( + settings.min_batch_size, + math.ceil( + settings.target_groups_per_step + * self.config.freshness_min_batch_floor_fraction + ), + ), "queue_maxsize": recommended_queue_size( target_groups_per_step=settings.target_groups_per_step, limit_steps_off_policy=policy_age_limit_steps, num_rollout_workers=settings.num_rollout_workers, running_reserve_fraction=self.config.queue_running_reserve_fraction, - ) + ), } ) diff --git a/src/art/pipeline_tuner/autotune.py b/src/art/pipeline_tuner/autotune.py index aa89d0471..cf2708282 100644 --- a/src/art/pipeline_tuner/autotune.py +++ b/src/art/pipeline_tuner/autotune.py @@ -44,7 +44,6 @@ def _ceil_to_multiple(value: float, multiple: int, *, minimum: int = 1) -> int: _VLLM_SCRAPE_GROUP_TOLERANCE_S = 0.05 -_TRAINER_PADDING_EPSILON = 1e-9 class PackingProjection(pydantic.BaseModel): @@ -58,14 +57,6 @@ class PackingOutcome(pydantic.BaseModel): packed_sequences: int = pydantic.Field(ge=1) -def _trainer_underfeed_score(*, idle_frac: float, padding_ratio: float) -> float: - denominator = max( - _TRAINER_PADDING_EPSILON, - 1.0 + _TRAINER_PADDING_EPSILON - max(0.0, min(1.0, padding_ratio)), - ) - return max(0.0, idle_frac) / denominator - - class PipelineAutotuner: def __init__( self, @@ -98,6 +89,9 @@ def __init__( self._target_candidate: int | None = None self._target_candidate_count = 0 self._stale_backlog_active = False + self._min_batch_trial_baseline_collect_s: float | None = None + self._min_batch_trial_batch_size: int | None = None + self._min_batch_trial_failed_windows = 0 self._emitted_recommendations: set[str] = set() def on_metric(self, rec: PipelineMetric) -> TunerDecision | None: @@ -221,10 +215,8 @@ def step_values(name: str) -> list[float]: end_step=window_steps[-1], window_start_s=t0, window_end_s=t1, - trainer_underfeed_score=_trainer_underfeed_score( - idle_frac=trainer_idle_frac, - padding_ratio=padding_ratio_mean, - ), + collect_batch_s=collect / len(window_steps), + trainer_underfeed_score=max(0.0, trainer_idle_frac), vllm_pressure=_vllm_pressure( vllm_metrics, window_start_s=t0, window_end_s=t1 ), @@ -314,6 +306,8 @@ def _decide(self, stats: TunerWindowStats) -> TunerDecision: target_changed = ( updated.target_groups_per_step != previous.target_groups_per_step ) + if target_changed: + self._clear_min_batch_trial() stale_backlog_active = self._update_stale_backlog_state(stats) action = "hold" reason = "inside hysteresis band or already balanced" @@ -361,7 +355,7 @@ def _decide(self, stats: TunerWindowStats) -> TunerDecision: reason = "both sides are loaded; no throughput-safe online change" if not target_changed and not stale_backlog_active: - min_update = self._min_batch_adjustment(updated, stats, state, action) + min_update = self._min_batch_adjustment(updated, stats, action) if min_update is not None: updated, action, reason = min_update @@ -397,9 +391,38 @@ def _min_batch_adjustment( self, settings: PipelineTuneSettings, stats: TunerWindowStats, - state: str, action: str, ) -> tuple[PipelineTuneSettings, str, str] | None: + if self._min_batch_trial_baseline_collect_s is not None: + if settings.min_batch_size != self._min_batch_trial_batch_size: + self._clear_min_batch_trial() + elif ( + stats.trainer_underfeed_score + <= self.config.trainer_min_batch_raise_score + ): + self._clear_min_batch_trial() + return self._raise_min_batch( + settings, + "trainer collection idle fell below the min-batch threshold", + ) + elif stats.collect_batch_s >= ( + self._min_batch_trial_baseline_collect_s + * self.config.min_batch_collect_improvement_ratio + ): + self._min_batch_trial_failed_windows += 1 + if ( + self._min_batch_trial_failed_windows + >= self.config.min_batch_trial_windows + ): + self._clear_min_batch_trial() + return self._raise_min_batch( + settings, + "smaller batches did not reduce collection time enough", + ) + return None + else: + self._clear_min_batch_trial() + if ( action != "increase_workers" and stats.trainer_underfeed_score @@ -414,6 +437,9 @@ def _min_batch_adjustment( ) new_min = max(floor, round(settings.min_batch_size * 0.85)) if new_min < settings.min_batch_size: + self._min_batch_trial_baseline_collect_s = stats.collect_batch_s + self._min_batch_trial_batch_size = new_min + self._min_batch_trial_failed_windows = 0 return ( settings.model_copy( update={"min_batch_size": min(new_min, settings.max_batch_size)} @@ -421,23 +447,35 @@ def _min_batch_adjustment( "lower_min_batch_size", "trainer is severely underfed and rollout workers are not being increased", ) - should_raise = action == "decrease_workers" or state in { - "inference_under_train_over", - "inference_balanced_train_over", - } - if should_raise and settings.min_batch_size < settings.max_batch_size: - new_min = min( - settings.max_batch_size, - max(settings.min_batch_size + 1, round(settings.min_batch_size * 1.15)), + if stats.trainer_underfeed_score <= self.config.trainer_min_batch_raise_score: + return self._raise_min_batch( + settings, + "trainer collection idle is low enough to use denser batches", ) - if new_min > settings.min_batch_size: - return ( - settings.model_copy(update={"min_batch_size": new_min}), - "raise_min_batch_size", - "trainer is saturated enough to use denser batches before reducing workers", - ) return None + def _raise_min_batch( + self, + settings: PipelineTuneSettings, + reason: str, + ) -> tuple[PipelineTuneSettings, str, str] | None: + if settings.min_batch_size >= settings.max_batch_size: + return None + new_min = min( + settings.max_batch_size, + max(settings.min_batch_size + 1, round(settings.min_batch_size * 1.15)), + ) + return ( + settings.model_copy(update={"min_batch_size": new_min}), + "raise_min_batch_size", + reason, + ) + + def _clear_min_batch_trial(self) -> None: + self._min_batch_trial_baseline_collect_s = None + self._min_batch_trial_batch_size = None + self._min_batch_trial_failed_windows = 0 + def _emit_stable_recommendations(self, decision: TunerDecision) -> None: recommendations = self._stable_recommendations() decision.recommendations.extend(message for _, message in recommendations) @@ -539,10 +577,11 @@ def _settings_with_recomputed_queue( if adapt_target else settings.target_groups_per_step ) - min_batch = min(settings.min_batch_size, target) + floor = math.ceil(target * self.config.freshness_min_batch_floor_fraction) + min_batch = max(floor, min(settings.min_batch_size, target)) if adapt_target and target > settings.target_groups_per_step: ratio = settings.min_batch_size / max(1, settings.max_batch_size) - min_batch = min(target, max(1, round(target * ratio))) + min_batch = max(floor, min(target, max(1, round(target * ratio)))) # Packed sequence length is the user's cap on target/max batch size. If a # run should never use larger train batches, lower packed_sequence_length. queue = recommended_queue_size( @@ -829,6 +868,10 @@ def build_initial_settings( int(config.initial_min_groups_per_packed_sequence) * target_slots, max_batch, ) + min_batch = max( + min_batch, + math.ceil(max_batch * config.freshness_min_batch_floor_fraction), + ) queue = recommended_queue_size( target_groups_per_step=max_batch, limit_steps_off_policy=policy_age_limit_steps, diff --git a/src/art/pipeline_tuner/config.py b/src/art/pipeline_tuner/config.py index 472de2325..e17eabba6 100644 --- a/src/art/pipeline_tuner/config.py +++ b/src/art/pipeline_tuner/config.py @@ -60,10 +60,15 @@ class PipelineAutotuneConfig(pydantic.BaseModel): stale_clear_frac: float = pydantic.Field(default=0.10, ge=0.0, le=1.0) padding_high_frac: float = pydantic.Field(default=0.25, ge=0.0, le=1.0) trainer_min_batch_lower_score: float = pydantic.Field(default=0.15, ge=0.0) + trainer_min_batch_raise_score: float = pydantic.Field(default=0.10, ge=0.0) + min_batch_collect_improvement_ratio: float = pydantic.Field( + default=0.85, gt=0.0, le=1.0 + ) + min_batch_trial_windows: int = pydantic.Field(default=2, ge=1) recommendation_min_windows: int = pydantic.Field(default=5, ge=1) recommendation_consecutive_holds: int = pydantic.Field(default=2, ge=1) freshness_min_batch_floor_fraction: float = pydantic.Field( - default=0.50, gt=0.0, le=1.0 + default=0.85, gt=0.0, le=1.0 ) target_group_change_windows: int = pydantic.Field(default=1, ge=1) target_group_increase_fraction: float = pydantic.Field(default=0.25, gt=0.0, le=1.0) @@ -83,6 +88,10 @@ class PipelineAutotuneConfig(pydantic.BaseModel): def validate_stale_hysteresis(self) -> "PipelineAutotuneConfig": if self.stale_clear_frac > self.stale_high_frac: raise ValueError("stale_clear_frac must be <= stale_high_frac") + if self.trainer_min_batch_raise_score > self.trainer_min_batch_lower_score: + raise ValueError( + "trainer_min_batch_raise_score must be <= trainer_min_batch_lower_score" + ) return self @@ -130,6 +139,7 @@ class TunerWindowStats(pydantic.BaseModel): end_step: int window_start_s: float = 0.0 window_end_s: float = 0.0 + collect_batch_s: float = 0.0 trainer_underfeed_score: float = 0.0 vllm_pressure: float = 0.0 queue_put_wait_frac: float = 0.0 From 75f3193e97c82f94783ce78b040620965d8e8455 Mon Sep 17 00:00:00 2001 From: FurtherAI Date: Fri, 17 Jul 2026 08:20:30 +0000 Subject: [PATCH 5/7] Report CP-compacted Megatron token throughput --- src/art/megatron/train.py | 36 ++++++++++++++++++++++++++++++++---- src/art/metrics.py | 14 +++++++++++++- src/art/model.py | 1 + 3 files changed, 46 insertions(+), 5 deletions(-) diff --git a/src/art/megatron/train.py b/src/art/megatron/train.py index 021329247..284620115 100644 --- a/src/art/megatron/train.py +++ b/src/art/megatron/train.py @@ -108,6 +108,7 @@ _zero_contribution_inputs, _zero_contribution_sft_inputs, build_micro_sample_indices, + build_micro_sample_indices_by_dp_rank, build_rl_hybridep_token_counts, build_sft_hybridep_token_counts, resolve_global_grad_accumulation_sequences, @@ -649,7 +650,12 @@ def run_megatron_rl_job( hybridep_token_counts=hybridep_token_counts, ) train_step_s = time.perf_counter() - train_step_started - global_packed_train_tokens = _global_packed_train_tokens(micro_inputs) + global_packed_train_tokens = _global_packed_train_tokens( + packed_tensors, + step_index=step_index, + num_sequences=num_sequences, + global_grad_accumulation_sequences=global_grad_accumulation_sequences, + ) print0( runtime.rank, "Correlation between old and new probabilities:", @@ -1108,6 +1114,7 @@ def _log_rl_step_result( "loss/grad_norm": step_result.grad_norm, "loss/probs_corr": step_result.probs_corr, TRAIN_GRADIENT_STEPS_KEY: num_gradient_steps, + "data/step_executed_packed_train_tokens": packed_train_tokens, "throughput/train_packed_tok_per_s": train_packed_tok_per_s, } if step_result.kl_policy_ref is not None: @@ -1118,9 +1125,30 @@ def _log_rl_step_result( log_file.write(log_msg + "\n") -def _global_packed_train_tokens(micro_inputs: list[PackedTensors]) -> int: - local_tokens = sum(int(micro["tokens"].numel()) for micro in micro_inputs) - return local_tokens * ps.get_data_parallel_world_size(with_context_parallel=False) +def _global_packed_train_tokens( + packed_tensors: PackedTensors, + *, + step_index: int, + num_sequences: int, + global_grad_accumulation_sequences: int | None, +) -> int: + sample_rows = build_micro_sample_indices_by_dp_rank( + step_index=step_index, + num_sequences=num_sequences, + global_grad_accumulation_sequences=global_grad_accumulation_sequences, + ) + sequence_length = int(packed_tensors["tokens"].shape[1]) + if ps.get_context_parallel_world_size() <= 1: + return sum(len(row) for row in sample_rows) * sequence_length + return sum( + int( + (packed_tensors["group_ids"][0 if index is None else index] != -1) + .sum() + .item() + ) + for row in sample_rows + for index in row + ) def _save_optimizer( diff --git a/src/art/metrics.py b/src/art/metrics.py index 391e3499f..f7d8ccdb5 100644 --- a/src/art/metrics.py +++ b/src/art/metrics.py @@ -113,11 +113,23 @@ class MetricDefinition(pydantic.BaseModel): higher_is_better=False, dashboard_default=True, ), + MetricDefinition( + key="data/step_executed_packed_train_tokens", + title="Megatron executed packed train tokens", + description=( + "packed token rows included in Megatron throughput; CP excludes " + "configured packed-row padding that is not dispatched" + ), + kind="counter", + unit="tokens", + higher_is_better=None, + ), MetricDefinition( key="throughput/train_packed_tok_per_s", title="Megatron packed train tokens per second", description=( - "physical packed training-token throughput reported by the Megatron worker" + "physical training-token throughput reported by the Megatron worker; " + "CP excludes configured packed-row padding that is not dispatched" ), kind="rate", unit="tok/s", diff --git a/src/art/model.py b/src/art/model.py index 32de0a065..1dd47cbb4 100644 --- a/src/art/model.py +++ b/src/art/model.py @@ -340,6 +340,7 @@ def __getattr__(self, name: str) -> Any: "offpolicy/token_weighted_policy_age_p95_steps", "throughput/accepted_train_tok_per_s", "throughput/train_packed_tok_per_s", + "data/step_executed_packed_train_tokens", "data/step_trainable_assistant_tokens", "data/step_non_padding_train_tokens", "data/step_padding_ratio", From 5066feb23cb5d684e1e2b8c54d9551ecc79a71fc Mon Sep 17 00:00:00 2001 From: FurtherAI Date: Fri, 17 Jul 2026 08:50:36 +0000 Subject: [PATCH 6/7] Preserve vLLM pressure gate for min-batch raises --- src/art/pipeline_tuner/autotune.py | 20 +++++++++++++++++--- 1 file changed, 17 insertions(+), 3 deletions(-) diff --git a/src/art/pipeline_tuner/autotune.py b/src/art/pipeline_tuner/autotune.py index cf2708282..5429478be 100644 --- a/src/art/pipeline_tuner/autotune.py +++ b/src/art/pipeline_tuner/autotune.py @@ -355,7 +355,12 @@ def _decide(self, stats: TunerWindowStats) -> TunerDecision: reason = "both sides are loaded; no throughput-safe online change" if not target_changed and not stale_backlog_active: - min_update = self._min_batch_adjustment(updated, stats, action) + min_update = self._min_batch_adjustment( + updated, + stats, + action, + inference_over=inference_over, + ) if min_update is not None: updated, action, reason = min_update @@ -392,12 +397,15 @@ def _min_batch_adjustment( settings: PipelineTuneSettings, stats: TunerWindowStats, action: str, + *, + inference_over: bool, ) -> tuple[PipelineTuneSettings, str, str] | None: if self._min_batch_trial_baseline_collect_s is not None: if settings.min_batch_size != self._min_batch_trial_batch_size: self._clear_min_batch_trial() elif ( - stats.trainer_underfeed_score + not inference_over + and stats.trainer_underfeed_score <= self.config.trainer_min_batch_raise_score ): self._clear_min_batch_trial() @@ -409,6 +417,8 @@ def _min_batch_adjustment( self._min_batch_trial_baseline_collect_s * self.config.min_batch_collect_improvement_ratio ): + if inference_over: + return None self._min_batch_trial_failed_windows += 1 if ( self._min_batch_trial_failed_windows @@ -447,7 +457,11 @@ def _min_batch_adjustment( "lower_min_batch_size", "trainer is severely underfed and rollout workers are not being increased", ) - if stats.trainer_underfeed_score <= self.config.trainer_min_batch_raise_score: + if ( + not inference_over + and stats.trainer_underfeed_score + <= self.config.trainer_min_batch_raise_score + ): return self._raise_min_batch( settings, "trainer collection idle is low enough to use denser batches", From 182187c93f9e20f677b98f0e34cd675687accdfc Mon Sep 17 00:00:00 2001 From: FurtherAI Date: Fri, 17 Jul 2026 08:50:51 +0000 Subject: [PATCH 7/7] Make pipeline shutdown event driven --- src/art/pipeline_trainer/trainer.py | 98 +++++++++++++++-------------- 1 file changed, 51 insertions(+), 47 deletions(-) diff --git a/src/art/pipeline_trainer/trainer.py b/src/art/pipeline_trainer/trainer.py index f9bbde919..3215f93d3 100644 --- a/src/art/pipeline_trainer/trainer.py +++ b/src/art/pipeline_trainer/trainer.py @@ -2,7 +2,7 @@ import asyncio from collections import Counter -from collections.abc import Mapping +from collections.abc import Awaitable, Mapping from contextlib import AsyncExitStack, asynccontextmanager from datetime import datetime, timezone import inspect @@ -282,6 +282,7 @@ def __init__( self._scheduled_eval_leases: dict[int, AsyncExitStack] = {} self.state = PipelineState() + self._stop_event = asyncio.Event() self._scenario_lock = asyncio.Lock() self._scenario_iter: AsyncIterator[ScenarioT] | None = _to_async_iterator( scenarios @@ -466,28 +467,23 @@ def _sync_signal_handler(signum: int, _frame: object | None) -> None: def request_stop(self) -> None: """Request a clean shutdown of the pipeline stages.""" - if self.state.done: - return self.state.done = True + self._stop_event.set() - async def _notify_policy() -> None: - async with self.state.policy_updated: - self.state.policy_updated.notify_all() - + async def _await_or_stop(self, awaitable: Awaitable[T]) -> tuple[bool, T | None]: + operation = asyncio.ensure_future(awaitable) + stop_wait = asyncio.create_task(self._stop_event.wait()) + tasks = (operation, stop_wait) try: - loop = asyncio.get_running_loop() - except RuntimeError: - loop = None - - if loop is None: - return - - loop.create_task(_notify_policy()) - if self._output_queue is not None: - try: - self._output_queue.put_nowait(None) - except asyncio.QueueFull: - loop.create_task(self._output_queue.put(None)) + done, _ = await asyncio.wait(tasks, return_when=asyncio.FIRST_COMPLETED) + if operation in done: + return True, operation.result() + return False, None + finally: + for task in tasks: + if not task.done(): + task.cancel() + await asyncio.gather(*tasks, return_exceptions=True) async def _finalize_backend_training(self) -> None: if not self._backend_training_completed: @@ -677,13 +673,17 @@ async def _get_next_scenario(self) -> ScenarioT | None: if self._scenario_source_exhausted: return None try: - scenario = await anext(self._scenario_iter) + completed, scenario = await self._await_or_stop( + anext(self._scenario_iter) + ) except StopAsyncIteration: self._scenario_source_exhausted = True return None + if not completed: + return None self.state.scenario_offset += 1 self.state.total_scenarios_consumed += 1 - return scenario + return cast(ScenarioT, scenario) async def _wait_for_policy(self) -> None: if self.max_steps_off_policy is None: @@ -694,7 +694,11 @@ async def _wait_for_policy(self) -> None: and self.state.policy_version < self.state.next_training_step - self.max_steps_off_policy ): - await self.state.policy_updated.wait() + completed, _ = await self._await_or_stop( + self.state.policy_updated.wait() + ) + if not completed: + return @asynccontextmanager async def _checkpoint_lease(self, step: int) -> AsyncIterator[None]: @@ -854,10 +858,7 @@ async def _rollout_stage(self) -> None: and self._output_queue is not None ): print("Scenario source exhausted; draining completed rollouts.") - try: - self._output_queue.put_nowait(None) - except asyncio.QueueFull: - await self._output_queue.put(None) + await self._await_or_stop(self._output_queue.put(None)) async def _training_stage(self) -> None: if self._output_queue is None: @@ -868,10 +869,8 @@ async def _training_stage(self) -> None: current_step + self.max_steps if self.max_steps is not None else None ) if stop_at_step is not None and current_step >= stop_at_step: - self.state.done = True self._persist_state(current_step) - async with self.state.policy_updated: - self.state.policy_updated.notify_all() + self.request_stop() return stop_after_batch = False @@ -1049,10 +1048,8 @@ async def _training_stage(self) -> None: if stop_after_batch: break - self.state.done = True self._persist_state(current_step) - async with self.state.policy_updated: - self.state.policy_updated.notify_all() + self.request_stop() async def _collect_batch( self, current_step: int @@ -1063,7 +1060,10 @@ async def _collect_batch( saw_sentinel = False while len(batch) < self.min_batch_size: - item = await self._output_queue.get() + completed, item = await self._await_or_stop(self._output_queue.get()) + if not completed: + saw_sentinel = True + break if item is None: saw_sentinel = True break @@ -1113,11 +1113,16 @@ async def _eval_stage(self) -> None: return pending_eval: asyncio.Task[None] | None = None - while not self.state.done or not self._eval_queue.empty(): + while True: try: - step = await asyncio.wait_for(self._eval_queue.get(), timeout=1.0) - except asyncio.TimeoutError: - continue + step = self._eval_queue.get_nowait() + except asyncio.QueueEmpty: + if self.state.done: + break + completed, step = await self._await_or_stop(self._eval_queue.get()) + if not completed: + continue + assert step is not None if pending_eval is not None and not pending_eval.done(): try: @@ -1137,7 +1142,10 @@ async def _status_loop(self) -> None: sleep_seconds = min(1.0, max(0.2, self._status_log_interval_seconds / 10)) while not self.state.done: self._status.log_if_due() - await asyncio.sleep(sleep_seconds) + try: + await asyncio.wait_for(self._stop_event.wait(), timeout=sleep_seconds) + except asyncio.TimeoutError: + continue async def _run_eval(self, step: int) -> None: assert self.eval_fn is not None @@ -1310,7 +1318,7 @@ def _trigger_collapse(self) -> None: if self._collapse_triggered: return self._collapse_triggered = True - self.state.done = True + self.request_stop() print( "\n" "========================================\n" @@ -1819,13 +1827,9 @@ def _is_scalar_metadata(value: object) -> bool: async def _put_output_group(self, group: TrajectoryGroup) -> float: assert self._output_queue is not None queue_wait_started = time.monotonic() - while not self.state.done: - try: - await asyncio.wait_for(self._output_queue.put(group), timeout=1.0) - self._status.note_group_enqueued(group) - return time.monotonic() - queue_wait_started - except asyncio.TimeoutError: - continue + completed, _ = await self._await_or_stop(self._output_queue.put(group)) + if completed: + self._status.note_group_enqueued(group) return time.monotonic() - queue_wait_started def _record_producer_rollout_timings(