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
7 changes: 7 additions & 0 deletions nemo_rl/environments/nemo_gym.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
16 changes: 10 additions & 6 deletions tests/unit/environments/test_nemo_gym_router_replay.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,20 +27,24 @@ 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": [
{
"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,
},
]
},
Expand All @@ -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():
Expand Down
Loading