Skip to content
Merged
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 nemo_rl/algorithms/single_controller.py
Original file line number Diff line number Diff line change
Expand Up @@ -2183,14 +2183,18 @@ async def _train_pump(self) -> None:
self._logger.log_metrics(
step_metrics, step=self._train_steps, prefix="train"
)
# Must precede the step_finished=True log below. That log commits
# the wandb step, and wandb silently discards anything logged
# against a step it has already committed -- no exception, no
# failed return, just an empty chart. grpo_sync had the same bug.
self._log_data_plane_metrics(total_time)
# step_finished=True here since this is the final log of our current step.
self._logger.log_metrics(
timing_metrics,
step=self._train_steps,
prefix="timing/train",
step_finished=True,
)
self._log_data_plane_metrics(total_time)
self._timer.reset()

# min sample version refers to the version each consumed sample was
Expand Down
45 changes: 45 additions & 0 deletions tests/unit/data_plane/test_observability.py
Original file line number Diff line number Diff line change
Expand Up @@ -1263,3 +1263,48 @@ def test_inspection_snapshot_does_not_steal_codec_time():
0.010, rel=1e-3
)
client.close()


def test_no_algorithm_logs_data_plane_after_committing_the_step():
"""``log_metrics(..., step_finished=True)`` commits the wandb step, and
wandb discards anything logged against a step it has already committed --
without raising, and without a falsy return to check.

This has bitten twice. ``grpo_sync`` was caught only because a real run
showed 85 logged keys and zero ``data_plane/*``. ``single_controller``
carried the same code -- its docstring says it mirrors ``grpo_sync`` --
and so carried the same bug, unnoticed, because no run exercised it.

Asserting the invariant for every algorithm rather than for one call site
is the point: a third wiring would otherwise repeat it. Source order is
the only observable, because the drop happens inside wandb where a fake
logger sees a perfectly ordinary call.
"""
import pathlib

import nemo_rl

algorithms = pathlib.Path(nemo_rl.__file__).parent / "algorithms"
checked = []
for path in sorted(algorithms.glob("*.py")):
# Comments mention the flag too, so they cannot be part of the search.
source = "\n".join(
line
for line in path.read_text().splitlines()
if not line.lstrip().startswith("#")
)
if "_log_data_plane_metrics(" not in source:
continue
if "step_finished=True" not in source:
continue
checked.append(path.name)
# rindex, not index: the first occurrence is the *definition*, which
# naturally precedes everything. The call site is what has to come
# before the commit, and it is the last occurrence.
assert source.rindex("_log_data_plane_metrics(") < source.index(
"step_finished=True"
), (
f"{path.name}: data-plane metrics are logged after the "
"step_finished=True commit, so wandb will discard them"
)
assert len(checked) >= 2, f"expected sync and single-controller, got {checked}"
Loading