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
6 changes: 5 additions & 1 deletion docs/en/get_started/customization.md
Original file line number Diff line number Diff line change
Expand Up @@ -170,6 +170,11 @@ def filter_function(args, samples: list[Sample]) -> None

**Note**: This function should directly modify the `remove_sample` attribute of each `Sample` object.

The default sample-to-training-data converter propagates this decision as
`removed_by_filter`. A custom `--custom-convert-samples-to-train-data-path`
that emits an all-zero loss mask must also emit a parallel boolean
`removed_by_filter` field. The trainer rejects unexplained all-zero masks.

**Use Cases**:
- Filtering samples based on response quality
- Implementing selective training strategies
Expand Down Expand Up @@ -439,4 +444,3 @@ For detailed explanation of R3 and MilesRouter, see [Miles Router](../advanced/m
def custom_model_provider(pre_process: bool, post_process: bool, vp_stage: int | None = None) -> GPTModel
```


9 changes: 7 additions & 2 deletions miles/backends/experimental/fsdp_utils/actor.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,7 +21,7 @@
from miles.utils.tracking_utils import init_tracking

from ....utils.profile_utils import TrainProfiler
from ...training_utils.ci_utils import check_grad_norm
from ...training_utils.ci_utils import assert_rollout_engine_weight_versions, check_grad_norm
from ...training_utils.data import DataIterator, get_batch, get_data_iterator, get_rollout_data
from ...training_utils.log_utils import (
aggregate_forward_results,
Expand Down Expand Up @@ -569,7 +569,12 @@ def update_weights(self) -> None: # type: ignore[override]

self.weight_updater.update_weights()

if self.args.ci_test and len(rollout_engines) > 0:
if getattr(self.args, "check_all_engine_weight_versions", False):
assert_rollout_engine_weight_versions(
rollout_engines,
self.weight_updater.weight_version,
)
elif self.args.ci_test and len(rollout_engines) > 0:
engine = random.choice(rollout_engines)
engine_version = ray.get(engine.get_weight_version.remote())
if str(engine_version) != str(self.weight_updater.weight_version):
Expand Down
88 changes: 71 additions & 17 deletions miles/backends/megatron_utils/actor.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,7 @@
from ...utils.profile_utils import TrainProfiler
from ...utils.tensor_backper import TensorBackuper
from ..training_utils.cp_utils import slice_with_cp
from ..training_utils.ci_utils import assert_rollout_engine_weight_versions
from ..training_utils.data import DataIterator, get_data_iterator, get_rollout_data, sync_actor_critic_data
from ..training_utils.log_utils import log_cpu_memory, log_perf_data, log_rollout_data
from ..training_utils.loss import compute_advantages_and_returns, get_log_probs_and_entropy, get_values
Expand Down Expand Up @@ -190,6 +191,13 @@ def _sum_float(x):
return float(sum(float(v) for v in x))
return float(x)

def _all_zero(x):
if torch.is_tensor(x):
return bool(torch.all(x == 0).item())
if isinstance(x, (list, tuple)):
return all(float(v) == 0 for v in x)
return float(x) == 0

def _summarize_vector_list(key, limit=3):
if not _present(key):
return f"{key}=MISSING"
Expand Down Expand Up @@ -227,6 +235,7 @@ def _basic_batch_summary():
"rewards",
"response_lengths",
"total_lengths",
"removed_by_filter",
"loss_masks",
"tokens",
"input_ids",
Expand All @@ -241,6 +250,16 @@ def _basic_batch_summary():
lines.append(_summarize_vector_list(key))

# Numeric aggregate summary.
try:
if _present("removed_by_filter") and _is_seq(rollout_data["removed_by_filter"]):
flags = rollout_data["removed_by_filter"]
lines.append(
f"removed_by_filter: count={len(flags)} "
f"removed={sum(bool(flag) for flag in flags)}"
)
except Exception as e:
lines.append(f"removed_by_filter aggregate failed: {type(e).__name__}: {e}")

try:
if _present("response_lengths") and _is_seq(rollout_data["response_lengths"]):
rs = [int(x) for x in rollout_data["response_lengths"]]
Expand Down Expand Up @@ -328,6 +347,28 @@ def _basic_batch_summary():
if got != n:
_add_error(f"{key!r} length mismatch: got {got}, expected {n}")

removed_by_filter_flags = [False] * n
if _present("removed_by_filter"):
candidate_flags = rollout_data["removed_by_filter"]
candidate_flags_valid = True
if not _is_seq(candidate_flags):
_add_error(
f"'removed_by_filter' must be list/tuple, got {type(candidate_flags).__name__}"
)
candidate_flags_valid = False
elif len(candidate_flags) != n:
_add_error(f"'removed_by_filter' length mismatch: got {len(candidate_flags)}, expected {n}")
candidate_flags_valid = False
else:
for i, removed in enumerate(candidate_flags):
if not isinstance(removed, bool):
_add_error(
f"removed_by_filter[{i}] must be bool, got {type(removed).__name__}"
)
candidate_flags_valid = False
if candidate_flags_valid:
removed_by_filter_flags = list(candidate_flags)

token_key = None
if _present("tokens"):
token_key = "tokens"
Expand Down Expand Up @@ -410,8 +451,24 @@ def _basic_batch_summary():
continue

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}")
removed_by_filter = removed_by_filter_flags[i]
mask_is_all_zero = _all_zero(mask)
if removed_by_filter:
if mask_is_all_zero:
_add_warning(
f"loss_masks[{i}] has no active tokens because removed_by_filter=True, "
f"sum={mask_sum}, response_len={resp}"
)
else:
_add_error(
f"loss_masks[{i}] is not all-zero despite removed_by_filter=True, "
f"sum={mask_sum}, response_len={resp}"
)
elif mask_sum <= 0:
_add_error(
f"loss_masks[{i}] has no active tokens without removed_by_filter=True, "
f"sum={mask_sum}, response_len={resp}"
)
if mask_sum > resp:
# Warning-only: float/weighted masks can legitimately have sum > resp.
_add_warning(
Expand Down Expand Up @@ -483,15 +540,6 @@ def check_vector_list(key, expected_lengths):
check_vector_list(key, response_lengths)
check_vector_list("rollout_log_probs", [] if cp_size > 1 else response_lengths)

# GRPO grouping diagnostics. Warning only because dynamic filtering can alter counts.
n_samples_per_prompt = int(getattr(args, "n_samples_per_prompt", 0) or 0)
if n_samples_per_prompt > 0 and n % n_samples_per_prompt != 0:
_add_warning(f"sample count {n} not divisible by n_samples_per_prompt={n_samples_per_prompt}")

grpo_group_size = int(getattr(args, "grpo_group_size", 0) or 0)
if grpo_group_size > 0 and n % grpo_group_size != 0:
_add_warning(f"sample count {n} not divisible by grpo_group_size={grpo_group_size}")

# This is important for your failure mode:
# If compute_advantages_and_returns will normalize, every rank that reaches
# it must have log_probs/values in the same structural state.
Expand Down Expand Up @@ -1013,13 +1061,19 @@ def update_weights(self) -> None:
self.weight_updater.update_weights()
print_memory("after update_weights")

if self.args.ci_test and len(rollout_engines) > 0 and not is_lora_enabled(self.args):
engine = random.choice(rollout_engines)
engine_version = ray.get(engine.get_weight_version.remote())
if str(engine_version) != str(self.weight_updater.weight_version):
raise RuntimeError(
f"Weight version mismatch! Engine: {engine_version}, Updater: {self.weight_updater.weight_version}"
if not is_lora_enabled(self.args):
if getattr(self.args, "check_all_engine_weight_versions", False):
assert_rollout_engine_weight_versions(
rollout_engines,
self.weight_updater.weight_version,
)
elif self.args.ci_test and len(rollout_engines) > 0:
engine = random.choice(rollout_engines)
engine_version = ray.get(engine.get_weight_version.remote())
if str(engine_version) != str(self.weight_updater.weight_version):
raise RuntimeError(
f"Weight version mismatch! Engine: {engine_version}, Updater: {self.weight_updater.weight_version}"
)

if getattr(self.args, "keep_old_actor", False):
if self.args.update_weights_interval == 1:
Expand Down
Loading
Loading