diff --git a/miles/backends/megatron_utils/actor.py b/miles/backends/megatron_utils/actor.py index c138ea87460..5bb4600c57a 100644 --- a/miles/backends/megatron_utils/actor.py +++ b/miles/backends/megatron_utils/actor.py @@ -411,7 +411,9 @@ def _basic_batch_summary(): mask_sum = _sum_float(mask) if mask_sum <= 0: - _add_error(f"loss_masks[{i}] has no active tokens, sum={mask_sum}, response_len={resp}") + # Warning-only: rollout sample filters zero a sample's loss_mask + # (remove_sample) to drop it from the gradient; the reducer clamps. + _add_warning(f"loss_masks[{i}] has no active tokens, sum={mask_sum}, response_len={resp}") if mask_sum > resp: # Warning-only: float/weighted masks can legitimately have sum > resp. _add_warning( diff --git a/miles/rollout/filter_hub/dynamic_sampling_filters.py b/miles/rollout/filter_hub/dynamic_sampling_filters.py index f39ff8c9c13..798c5d79546 100644 --- a/miles/rollout/filter_hub/dynamic_sampling_filters.py +++ b/miles/rollout/filter_hub/dynamic_sampling_filters.py @@ -6,10 +6,34 @@ __all__ = [ "check_reward_nonzero_std", "check_no_aborted", + "check_no_infra_failures", "drop_zero_std_groups_and_extreme_pass_rate", "drop_truncated_or_extreme_pass_rate", ] +# Harbor exit_status values (from _extract_exit_status, after _TIMEOUT_EXCEPTION_MAP) +# for infra / non-policy failures: the rollout was severed by the environment, +# engine, or verifier rather than the policy, so it isn't valid training signal. +INFRA_FAILURE_EXIT_STATUSES = frozenset( + { + "AgentTimeout", + "AgentTimeoutError", + "HealthcheckError", + "_K8sInternalInfraError", + "Cancelled", + "RewardFileNotFoundError", + "AgentSetupTimeout", + "AgentSetupTimeoutError", + "SqsConsumerError", + "VerifierTimeout", + "VerifierTimeoutError", + "EnvStartTimeout", + "EnvironmentStartTimeoutError", + "TimeoutError", + "AddTestsDirError", + } +) + def check_reward_nonzero_std(args, samples: list[Sample], **kwargs): rewards = [sample.get_reward_value(args) for sample in samples] @@ -46,6 +70,29 @@ def check_no_aborted(args, samples: list[Sample], **kwargs): return DynamicFilterOutput(keep=True) +def check_no_infra_failures(args, samples: list[Sample], **kwargs) -> DynamicFilterOutput: + """Reject the group if any rollout aborted or failed for an infra reason, + or if the group has zero reward std (all-correct / all-incorrect). + + Superset of ``check_no_aborted``: also drops samples whose Harbor + ``exit_status`` is an infra / non-policy failure. Keyed on ``exit_status`` + because such rollouts may still record COMPLETED turns; aborted samples + (zero records, no exit_status) are caught by the status check. Groups that + survive the infra checks are then passed through ``check_reward_nonzero_std`` + so all-same-reward groups (advantage 0, no gradient) are dropped too. + + --dynamic-sampling-filter-path miles.rollout.filter_hub.dynamic_sampling_filters.check_no_infra_failures + """ + flat_samples = list(_flatten_samples(samples)) + for sample in flat_samples: + if sample.status == Sample.Status.ABORTED: + return DynamicFilterOutput(keep=False, reason="group_has_aborted") + exit_status = (sample.metadata or {}).get("exit_status", "") + if exit_status in INFRA_FAILURE_EXIT_STATUSES: + return DynamicFilterOutput(keep=False, reason=f"group_has_{exit_status}") + return check_reward_nonzero_std(args, flat_samples, **kwargs) + + def drop_zero_std_groups_and_extreme_pass_rate(args, samples: list[Sample], **kwargs) -> DynamicFilterOutput: """Filter groups with near-zero reward std or extreme mean rewards. For 0/1 rewards, mean reward is equivalent to pass rate --- so this function can be used to filter diff --git a/miles/rollout/filter_hub/rollout_filters.py b/miles/rollout/filter_hub/rollout_filters.py index 72282d4c64b..30be747d712 100644 --- a/miles/rollout/filter_hub/rollout_filters.py +++ b/miles/rollout/filter_hub/rollout_filters.py @@ -1,7 +1,32 @@ from miles.rollout.filter_hub.dynamic_sampling_filters import _flatten_samples from miles.utils.types import Sample -__all__ = ["mask_truncated", "mask_truncated_and_llm_judge_failed"] +__all__ = ["mask_truncated", "mask_truncated_and_llm_judge_failed", "mask_token_truncated"] + +# Harbor exit_status values for rollouts truncated by running out of output +# tokens or context. Real attempts (verifier reward is meaningful) but the +# trajectory shouldn't get gradient — masked, reward kept in the group baseline. +TRUNCATION_EXIT_STATUSES = frozenset( + { + "BadRequestError", + "ContextWindowExceededError", + "OutputLengthExceededError", + } +) + + +def mask_token_truncated(args, samples: list[Sample]) -> None: + """Mask token/context-truncated samples (zero gradient, reward kept in baseline). + + Keyed on Harbor ``exit_status`` (BadRequest/ContextWindow record COMPLETED + turns, so a status check misses them); also covers TRUNCATED status. + + --rollout-sample-filter-path miles.rollout.filter_hub.rollout_filters.mask_token_truncated + """ + for sample in _flatten_samples(samples): + exit_status = (sample.metadata or {}).get("exit_status", "") + if exit_status in TRUNCATION_EXIT_STATUSES or sample.status == Sample.Status.TRUNCATED: + sample.remove_sample = True def mask_truncated(args, samples: list[Sample]) -> None: diff --git a/tests/fast/rollout/test_dynamic_sampling_filters.py b/tests/fast/rollout/test_dynamic_sampling_filters.py new file mode 100644 index 00000000000..cc855a16b7b --- /dev/null +++ b/tests/fast/rollout/test_dynamic_sampling_filters.py @@ -0,0 +1,77 @@ +from argparse import Namespace + +import pytest + +from miles.rollout.filter_hub.base_types import DynamicFilterOutput +from miles.rollout.filter_hub.dynamic_sampling_filters import ( + INFRA_FAILURE_EXIT_STATUSES, + check_no_infra_failures, +) +from miles.utils.types import Sample + + +ARGS = Namespace(reward_key=None) + + +def _sample( + reward: float | None, + *, + status: Sample.Status = Sample.Status.COMPLETED, + exit_status: str | None = None, +) -> Sample: + metadata = {} if exit_status is None else {"exit_status": exit_status} + return Sample(reward=reward, status=status, metadata=metadata) + + +@pytest.mark.parametrize("exit_status", sorted(INFRA_FAILURE_EXIT_STATUSES)) +def test_check_no_infra_failures_rejects_every_infra_exit_status(exit_status: str) -> None: + # The rewards are mixed, so this group would otherwise pass the std check. + samples = [_sample(0, exit_status=exit_status), _sample(1)] + + output = check_no_infra_failures(ARGS, samples) + + assert output == DynamicFilterOutput(keep=False, reason=f"group_has_{exit_status}") + + +def test_check_no_infra_failures_rejects_aborted_before_reward_check() -> None: + samples = [_sample(None, status=Sample.Status.ABORTED), _sample(1)] + + output = check_no_infra_failures(ARGS, samples) + + assert output == DynamicFilterOutput(keep=False, reason="group_has_aborted") + + +def test_check_no_infra_failures_keeps_nested_mixed_reward_group() -> None: + samples = [[_sample(0)], [_sample(1)]] + + output = check_no_infra_failures(ARGS, samples) + + assert bool(output.keep) + assert output.reason is None + + +@pytest.mark.parametrize( + ("reward", "expected_reason"), + [ + (1, "zero_std_1"), + (0, "zero_std_0"), + ], + ids=["all-correct", "all-incorrect"], +) +def test_check_no_infra_failures_rejects_zero_std_reward_groups( + reward: int, + expected_reason: str, +) -> None: + output = check_no_infra_failures(ARGS, [_sample(reward), _sample(reward)]) + + assert not bool(output.keep) + assert output.reason == expected_reason + + +def test_check_no_infra_failures_allows_policy_failure_with_mixed_rewards() -> None: + samples = [_sample(0, exit_status="TestsFailed"), _sample(1)] + + output = check_no_infra_failures(ARGS, samples) + + assert bool(output.keep) + assert output.reason is None