Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 3 additions & 1 deletion miles/backends/megatron_utils/actor.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
47 changes: 47 additions & 0 deletions miles/rollout/filter_hub/dynamic_sampling_filters.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]
Expand Down Expand Up @@ -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
Expand Down
27 changes: 26 additions & 1 deletion miles/rollout/filter_hub/rollout_filters.py
Original file line number Diff line number Diff line change
@@ -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:
Expand Down
77 changes: 77 additions & 0 deletions tests/fast/rollout/test_dynamic_sampling_filters.py
Original file line number Diff line number Diff line change
@@ -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
Loading