diff --git a/src/art/pipeline_trainer/binary_prefix_tool_pipeline.py b/src/art/pipeline_trainer/binary_prefix_tool_pipeline.py index 990af07fe..afd6ea3c4 100644 --- a/src/art/pipeline_trainer/binary_prefix_tool_pipeline.py +++ b/src/art/pipeline_trainer/binary_prefix_tool_pipeline.py @@ -336,8 +336,8 @@ async def scenario_iter(): num_rollout_workers=num_rollout_workers, min_batch_size=min_batch_size, max_batch_size=max_batch_size, - max_steps_off_policy=max_steps_off_policy, ), + max_steps_off_policy=max_steps_off_policy, learning_rate=float(os.environ.get("LEARNING_RATE", "1e-4")), log_interval_seconds=log_interval_seconds, eval_every_n_steps=eval_every_n_steps, diff --git a/src/art/pipeline_trainer/trainer.py b/src/art/pipeline_trainer/trainer.py index 99ca25868..f25f32fdc 100644 --- a/src/art/pipeline_trainer/trainer.py +++ b/src/art/pipeline_trainer/trainer.py @@ -152,7 +152,6 @@ def __init__( num_rollout_workers: int | None = None, min_batch_size: int | None = None, max_batch_size: int | None = None, - max_steps_off_policy: int | None = None, queue_maxsize: int | None = None, pipeline: PipelineRuntimeConfig | None = None, autotune: PipelineAutotuneConfig | None = None, @@ -167,6 +166,7 @@ def __init__( max_steps: int | None = None, # Discard handling discard_queue_multiplier: int = 100, + max_steps_off_policy: int | None = 4, limit_mean_steps_off_policy: float | None = None, score_reference_groups_per_step: float | None = None, score_reference_rollouts_per_group: float | None = None, @@ -191,7 +191,6 @@ def __init__( "num_rollout_workers": num_rollout_workers, "min_batch_size": min_batch_size, "max_batch_size": max_batch_size, - "max_steps_off_policy": max_steps_off_policy, "queue_maxsize": queue_maxsize, }.items() if value is not None @@ -223,6 +222,8 @@ def __init__( raise ValueError("log_interval_seconds must be > 0") if discard_queue_multiplier <= 0: raise ValueError("discard_queue_multiplier must be > 0") + if max_steps_off_policy is not None and max_steps_off_policy < 0: + raise ValueError("max_steps_off_policy must be >= 0") if limit_mean_steps_off_policy is not None and limit_mean_steps_off_policy < 0: raise ValueError("limit_mean_steps_off_policy must be >= 0") if checkpoint_retention_interval <= 0: @@ -246,7 +247,7 @@ def __init__( else 10 * pipeline.min_batch_size ) self.target_groups_per_step = self.max_batch_size - self.max_steps_off_policy = pipeline.max_steps_off_policy + self.max_steps_off_policy = max_steps_off_policy self.limit_mean_steps_off_policy = limit_mean_steps_off_policy self.queue_maxsize = pipeline.queue_maxsize self.learning_rate = learning_rate diff --git a/src/art/pipeline_tuner/config.py b/src/art/pipeline_tuner/config.py index c5bdf88b2..de1874fde 100644 --- a/src/art/pipeline_tuner/config.py +++ b/src/art/pipeline_tuner/config.py @@ -12,7 +12,6 @@ class PipelineRuntimeConfig(pydantic.BaseModel): num_rollout_workers: int = pydantic.Field(default=16, ge=1) min_batch_size: int = pydantic.Field(default=4, ge=1) max_batch_size: int | None = pydantic.Field(default=None, ge=1) - max_steps_off_policy: int | None = pydantic.Field(default=4, ge=0) queue_maxsize: int | None = pydantic.Field(default=None, ge=1) score_reference_groups_per_step: float | None = pydantic.Field(default=8.0, gt=0.0) score_reference_rollouts_per_group: float | None = pydantic.Field( diff --git a/tests/integration/megatron/trainability/test_live_length_trainability.py b/tests/integration/megatron/trainability/test_live_length_trainability.py index 5d970c73b..96ec424e6 100644 --- a/tests/integration/megatron/trainability/test_live_length_trainability.py +++ b/tests/integration/megatron/trainability/test_live_length_trainability.py @@ -829,8 +829,8 @@ async def rollout_fn( num_rollout_workers=rollout_workers, min_batch_size=1, max_batch_size=1, - max_steps_off_policy=max_steps_off_policy, ), + max_steps_off_policy=max_steps_off_policy, learning_rate=_get_env_float( "ART_MODEL_SUPPORT_LENGTH_LEARNING_RATE", _default_learning_rate(base_model), diff --git a/tests/unit/test_pipeline_trainer_metrics.py b/tests/unit/test_pipeline_trainer_metrics.py index f9d1d6802..48fe6d557 100644 --- a/tests/unit/test_pipeline_trainer_metrics.py +++ b/tests/unit/test_pipeline_trainer_metrics.py @@ -55,8 +55,8 @@ async def test_training_records_stale_and_zero_variance_discards( num_rollout_workers=1, min_batch_size=1, max_batch_size=1, - max_steps_off_policy=0, ), + max_steps_off_policy=0, eval_fn=None, max_steps=1, )