From e3ba8c2998894aa188d65b321f2d89ac8145549c Mon Sep 17 00:00:00 2001 From: alexchiu <7390474+zpqiu@users.noreply.github.com> Date: Mon, 3 Aug 2026 22:47:54 -0700 Subject: [PATCH 1/2] fix(nemo-gym): refresh final-token route across turns Signed-off-by: alexchiu <7390474+zpqiu@users.noreply.github.com> --- nemo_rl/environments/nemo_gym.py | 16 +++++++++++++ tests/unit/environments/test_nemo_gym.py | 29 ++++++++++++++++++++++++ 2 files changed, 45 insertions(+) diff --git a/nemo_rl/environments/nemo_gym.py b/nemo_rl/environments/nemo_gym.py index 30fb04954b2..69d7c055315 100644 --- a/nemo_rl/environments/nemo_gym.py +++ b/nemo_rl/environments/nemo_gym.py @@ -417,6 +417,22 @@ def _postprocess_nemo_gym_to_nemo_rl_result( "with enable_return_routed_experts." ) + # A generation response has no route for its final token because + # that token is never fed back into the model in the same request. + # pad_and_align_routed_expert_indices therefore leaves a valid dummy + # route at the end of every assistant turn. The next request does + # process that token as part of its prompt, so replace the previous + # turn's dummy with the newly captured real route before appending + # the new prompt suffix. + if routed_experts is not None and seen_token_ids: + previous_routes = nemo_rl_message_log[-1].get("routed_experts") + if not isinstance(previous_routes, torch.Tensor): + raise ValueError( + "A later NeMo Gym turn returned routed_experts but the " + "previous trainable turn did not carry routes to refresh." + ) + 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.py b/tests/unit/environments/test_nemo_gym.py index a4ca5759e98..e01582d5e5c 100644 --- a/tests/unit/environments/test_nemo_gym.py +++ b/tests/unit/environments/test_nemo_gym.py @@ -277,11 +277,25 @@ def batch_decode(self, batch): "prompt_token_ids": [1, 2], "generation_token_ids": [3], "generation_log_probs": [-0.1], + "routed_experts": [ + [[10, 11]], + [[20, 21]], + [[0, 1]], + ], }, { "prompt_token_ids": [1, 2, 3, 4, 5], "generation_token_ids": [6, 7], "generation_log_probs": [-0.2, -0.3], + "routed_experts": [ + [[10, 11]], + [[20, 21]], + [[30, 31]], + [[40, 41]], + [[50, 51]], + [[60, 61]], + [[0, 1]], + ], }, ] }, @@ -305,6 +319,21 @@ class _MockSelf: assert result["message_log"][1]["token_ids"].tolist() == [3] assert result["message_log"][2]["token_ids"].tolist() == [4, 5] assert result["message_log"][3]["token_ids"].tolist() == [6, 7] + assert result["message_log"][0]["routed_experts"].tolist() == [ + [[10, 11]], + [[20, 21]], + ] + # Turn 1 ended with the dummy [0, 1] route. Turn 2 re-processes token 3 + # in its prompt and must refresh that row to the real [30, 31] route. + assert result["message_log"][1]["routed_experts"].tolist() == [[[30, 31]]] + assert result["message_log"][2]["routed_experts"].tolist() == [ + [[40, 41]], + [[50, 51]], + ] + assert result["message_log"][3]["routed_experts"].tolist() == [ + [[60, 61]], + [[0, 1]], + ] assert nemo_gym_result["response"]["output"][0]["prompt_str"] == "1 2" assert nemo_gym_result["response"]["output"][0]["generation_str"] == "3" assert nemo_gym_result["response"]["output"][1]["prompt_str"] == "1 2 3 4 5" From 09a2e4017087f33e618d28ddf5260ade9d3858fb Mon Sep 17 00:00:00 2001 From: alexchiu <7390474+zpqiu@users.noreply.github.com> Date: Mon, 3 Aug 2026 23:18:02 -0700 Subject: [PATCH 2/2] fix(nemo-gym): simplify final-token route refresh Signed-off-by: alexchiu <7390474+zpqiu@users.noreply.github.com> --- nemo_rl/environments/nemo_gym.py | 17 +++-------- tests/unit/environments/test_nemo_gym.py | 29 ------------------- .../test_nemo_gym_router_replay.py | 16 ++++++---- 3 files changed, 14 insertions(+), 48 deletions(-) diff --git a/nemo_rl/environments/nemo_gym.py b/nemo_rl/environments/nemo_gym.py index 69d7c055315..1b3daba4612 100644 --- a/nemo_rl/environments/nemo_gym.py +++ b/nemo_rl/environments/nemo_gym.py @@ -417,21 +417,12 @@ def _postprocess_nemo_gym_to_nemo_rl_result( "with enable_return_routed_experts." ) - # A generation response has no route for its final token because - # that token is never fed back into the model in the same request. - # pad_and_align_routed_expert_indices therefore leaves a valid dummy - # route at the end of every assistant turn. The next request does - # process that token as part of its prompt, so replace the previous - # turn's dummy with the newly captured real route before appending - # the new prompt suffix. + # 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 not isinstance(previous_routes, torch.Tensor): - raise ValueError( - "A later NeMo Gym turn returned routed_experts but the " - "previous trainable turn did not carry routes to refresh." - ) - previous_routes[-1] = routed_experts[len(seen_token_ids) - 1] + 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) diff --git a/tests/unit/environments/test_nemo_gym.py b/tests/unit/environments/test_nemo_gym.py index e01582d5e5c..a4ca5759e98 100644 --- a/tests/unit/environments/test_nemo_gym.py +++ b/tests/unit/environments/test_nemo_gym.py @@ -277,25 +277,11 @@ def batch_decode(self, batch): "prompt_token_ids": [1, 2], "generation_token_ids": [3], "generation_log_probs": [-0.1], - "routed_experts": [ - [[10, 11]], - [[20, 21]], - [[0, 1]], - ], }, { "prompt_token_ids": [1, 2, 3, 4, 5], "generation_token_ids": [6, 7], "generation_log_probs": [-0.2, -0.3], - "routed_experts": [ - [[10, 11]], - [[20, 21]], - [[30, 31]], - [[40, 41]], - [[50, 51]], - [[60, 61]], - [[0, 1]], - ], }, ] }, @@ -319,21 +305,6 @@ class _MockSelf: assert result["message_log"][1]["token_ids"].tolist() == [3] assert result["message_log"][2]["token_ids"].tolist() == [4, 5] assert result["message_log"][3]["token_ids"].tolist() == [6, 7] - assert result["message_log"][0]["routed_experts"].tolist() == [ - [[10, 11]], - [[20, 21]], - ] - # Turn 1 ended with the dummy [0, 1] route. Turn 2 re-processes token 3 - # in its prompt and must refresh that row to the real [30, 31] route. - assert result["message_log"][1]["routed_experts"].tolist() == [[[30, 31]]] - assert result["message_log"][2]["routed_experts"].tolist() == [ - [[40, 41]], - [[50, 51]], - ] - assert result["message_log"][3]["routed_experts"].tolist() == [ - [[60, 61]], - [[0, 1]], - ] assert nemo_gym_result["response"]["output"][0]["prompt_str"] == "1 2" assert nemo_gym_result["response"]["output"][0]["generation_str"] == "3" assert nemo_gym_result["response"]["output"][1]["prompt_str"] == "1 2 3 4 5" 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():