diff --git a/nemo_rl/environments/nemo_gym.py b/nemo_rl/environments/nemo_gym.py index 30fb04954b2..1b3daba4612 100644 --- a/nemo_rl/environments/nemo_gym.py +++ b/nemo_rl/environments/nemo_gym.py @@ -417,6 +417,13 @@ def _postprocess_nemo_gym_to_nemo_rl_result( "with enable_return_routed_experts." ) + # The next prompt prefill supplies the real route for the previous + # turn's final token, whose decode route was padded. + if routed_experts is not None and seen_token_ids: + previous_routes = nemo_rl_message_log[-1].get("routed_experts") + if isinstance(previous_routes, torch.Tensor): + previous_routes[-1] = routed_experts[len(seen_token_ids) - 1] + prompt_start = len(seen_token_ids) prompt_end = len(prompt_token_ids) generation_start = prompt_end diff --git a/tests/unit/environments/test_nemo_gym_router_replay.py b/tests/unit/environments/test_nemo_gym_router_replay.py index d1171856213..15383d8af5b 100644 --- a/tests/unit/environments/test_nemo_gym_router_replay.py +++ b/tests/unit/environments/test_nemo_gym_router_replay.py @@ -27,6 +27,10 @@ def _routes(num_tokens: int) -> list[list[list[int]]]: def test_nemo_gym_postprocess_slices_routed_experts(): + first_turn_routes = _routes(3) + first_turn_routes[-1] = [[0, 1]] + second_turn_routes = _routes(7) + second_turn_routes[2] = [[30, 31]] nemo_gym_result = { "response": { "output": [ @@ -34,13 +38,13 @@ def test_nemo_gym_postprocess_slices_routed_experts(): "prompt_token_ids": [1, 2], "generation_token_ids": [3], "generation_log_probs": [-0.1], - "routed_experts": _routes(3), + "routed_experts": first_turn_routes, }, { "prompt_token_ids": [1, 2, 3, 4, 5], "generation_token_ids": [6, 7], "generation_log_probs": [-0.2, -0.3], - "routed_experts": _routes(7), + "routed_experts": second_turn_routes, }, ] }, @@ -58,13 +62,13 @@ class _MockSelf: message_log = result["message_log"] assert message_log[0]["token_ids"].tolist() == [1, 2] - assert message_log[0]["routed_experts"].tolist() == _routes(2) + assert message_log[0]["routed_experts"].tolist() == first_turn_routes[:2] assert message_log[1]["token_ids"].tolist() == [3] - assert message_log[1]["routed_experts"].tolist() == _routes(3)[2:3] + assert message_log[1]["routed_experts"].tolist() == second_turn_routes[2:3] assert message_log[2]["token_ids"].tolist() == [4, 5] - assert message_log[2]["routed_experts"].tolist() == _routes(7)[3:5] + assert message_log[2]["routed_experts"].tolist() == second_turn_routes[3:5] assert message_log[3]["token_ids"].tolist() == [6, 7] - assert message_log[3]["routed_experts"].tolist() == _routes(7)[5:7] + assert message_log[3]["routed_experts"].tolist() == second_turn_routes[5:7] def test_nemo_gym_postprocess_requires_routed_experts_when_configured():