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
39 changes: 33 additions & 6 deletions nemo_rl/environments/nemo_gym.py
Original file line number Diff line number Diff line change
Expand Up @@ -50,6 +50,12 @@
RolloutDataFailure,
http_status_is_infra,
)
from nemo_rl.experience.interfaces import (
NEMO_GYM_ATTEMPT_INDEX_KEY,
NEMO_GYM_CAPTURE_ID_KEY,
NEMO_GYM_ROLLOUT_ID_KEY,
nemo_gym_capture_key,
)
from nemo_rl.models.generation.interfaces import (
resolve_routed_experts_dtype_name_for_model,
should_use_async_rollouts,
Expand Down Expand Up @@ -231,12 +237,10 @@ class NemoGymConfig(TypedDict):
token_capture: NotRequired[Dict[str, Any] | None]


# Gym control-plane server name (the model server hosting the ledger) and the
# opaque run-body key rollout ids ride on (Gym's ROLLOUT_ID_KEY_NAME): the
# agent derives the id from the run body and stamps /ng-rollout/<id> on every
# model call, so the TQ sample id IS the capture key end to end.
# Gym control-plane server name for the model server hosting the capture ledger.
# Gym derives its physical capture key from the logical run-body rollout ID and
# numeric attempt index.
_POLICY_SERVER_NAME = "policy_model"
_NG_ROLLOUT_ID_BODY_KEY = "_ng_rollout_id"
_TOKEN_CAPTURE_CONTROL_PREFIX = "/training-token-capture/control"
_TOKEN_CAPTURE_CONTROL_ENV = "NEMO_GYM_TOKEN_CAPTURE_CONTROL_TOKEN"

Expand Down Expand Up @@ -602,6 +606,29 @@ async def run_rollouts(
)
tokenizer = self._tokenizer

for row in nemo_gym_examples:
logical_rollout_id = row.get(NEMO_GYM_ROLLOUT_ID_KEY)
attempt_index = row.get(NEMO_GYM_ATTEMPT_INDEX_KEY)
if not isinstance(logical_rollout_id, str) or not logical_rollout_id:
raise ValueError(
f"{NEMO_GYM_ROLLOUT_ID_KEY} must be a non-empty string"
)
if (
isinstance(attempt_index, bool)
or not isinstance(attempt_index, int)
or attempt_index < 0
):
raise ValueError(
f"{NEMO_GYM_ATTEMPT_INDEX_KEY} must be a non-negative integer"
)
expected_capture_id = nemo_gym_capture_key(
logical_rollout_id, attempt_index
)
if row.get(NEMO_GYM_CAPTURE_ID_KEY) != expected_capture_id:
raise ValueError(
f"{NEMO_GYM_CAPTURE_ID_KEY} must equal {expected_capture_id!r}"
)

from nemo_rl.utils.fastokens import maybe_patch_fastokens

maybe_patch_fastokens(bool(self.cfg.get("use_fastokens")))
Expand Down Expand Up @@ -710,7 +737,7 @@ async def _postprocess_receipt_mode(
assert isinstance(nemo_gym_result, dict), (
f"Hit a non-successful response when querying NeMo Gym for rollouts: {nemo_gym_result}"
)
rollout_id = nemo_gym_row[_NG_ROLLOUT_ID_BODY_KEY]
rollout_id = nemo_gym_row[NEMO_GYM_CAPTURE_ID_KEY]
# Gym's TERMINAL_RESPONSE_ID_KEY: the served response envelope id the
# harness kept (``response.id``), not the logical-request header.
terminal_response_id = nemo_gym_result.get("terminal_response_id")
Expand Down
12 changes: 12 additions & 0 deletions nemo_rl/experience/interfaces.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,9 @@
NEMO_GYM_GROUP_ID_KEY = "_ng_group_id"
NEMO_GYM_GROUP_ATTEMPT_KEY = "_ng_group_attempt"
NEMO_GYM_ROLLOUT_INDEX_KEY = "_ng_rollout_index"
NEMO_GYM_ROLLOUT_ID_KEY = "_ng_rollout_id"
NEMO_GYM_ATTEMPT_INDEX_KEY = "_ng_attempt_index"
NEMO_GYM_CAPTURE_ID_KEY = "_ng_capture_id"
NEXT_NEMO_GYM_TASK_INDEX_KEY = "next_ng_task_index"
# Unconsumed suffix of a gap-fill dataloader batch, carried in the async
# collector's rollouts state so a checkpoint cannot strand yielded prompts.
Expand All @@ -41,6 +44,15 @@
TRAINED_TASK_INDICES_KEY = "trained_task_indices"


def nemo_gym_capture_key(logical_rollout_id: str, attempt_index: int) -> str:
"""Return Gym's physical capture key for one logical rollout attempt."""
if attempt_index < 0:
raise ValueError("attempt_index must be non-negative")
if attempt_index == 0:
return logical_rollout_id
return f"{logical_rollout_id}-a{attempt_index}"


@dataclass
class Completion:
"""A single generated completion for one prompt."""
Expand Down
Loading
Loading