Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
15 commits
Select commit Hold shift + click to select a range
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
86 changes: 86 additions & 0 deletions nemo_gym/token_id_capture/adapters/megatron.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,86 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0

"""Dependency-light extraction adapter for Megatron Inference offloaded payloads."""

from __future__ import annotations

from collections.abc import Mapping, Sequence
from typing import Any


PREFIX_IDS_FIELD = "required_prefix_token_ids"
PROMPT_IDS_FIELD = "prompt_token_ids"
GENERATED_IDS_FIELD = "generated_token_ids"
GENERATED_LOGPROBS_FIELD = "generated_log_probs"


def _field(payload: Any, name: str) -> Any:
"""Read one field from a Megatron Inference payload object or an equivalent mapping."""
if isinstance(payload, Mapping):
return payload.get(name)
return getattr(payload, name, None)


def _sequence(payload: Any, name: str) -> Sequence[Any]:
value = _field(payload, name)
if value is None:
raise ValueError(f"Megatron offloaded payload carries no {name}")
if isinstance(value, (str, bytes)) or not isinstance(value, Sequence):
raise ValueError(f"Megatron offloaded payload field {name} must be a token-id sequence")
return value


def _token_ids(payload: Any, name: str) -> list[int]:
"""Read a token-id field, rejecting anything but plain integers.

Megatron Inference hands host-side ``list[int]`` values. A float, string, or
bool element means the payload is malformed; ``int()`` would silently
truncate or coerce it into a plausible-looking id.
"""
values = _sequence(payload, name)
if any(type(value) is not int for value in values):
raise ValueError(f"Megatron offloaded payload field {name} must contain only integer token ids")
return list(values)


def _log_probs(payload: Any, name: str) -> list[float]:
"""Read a log-probability field, rejecting non-numeric elements."""
values = _sequence(payload, name)
if any(isinstance(value, bool) or not isinstance(value, (int, float)) for value in values):
raise ValueError(f"Megatron offloaded payload field {name} must contain only numeric log probabilities")
return [float(value) for value in values]


class MegatronCaptureAdapter:
"""Translate Megatron inference request/response material at the framework boundary.

Megatron inference offloads exact prompt ids, generated ids, and selected-token log
probabilities as attributes on a payload object rather than as a chat
completion dict. The adapter reads either shape so extraction failures
flow through ``RolloutTokenCapture.complete_call_from_response`` and
poison the call with ``capture_failed`` coordinates, the same outcome the
vLLM adapter produces.
"""

def enter_prefix(self, request_payload: dict[str, Any], prefix_ids: list[int]) -> dict[str, Any]:
request_payload[PREFIX_IDS_FIELD] = list(prefix_ids)
return request_payload

def extract_prompt_ids(self, response_payload: Any) -> list[int]:
return _token_ids(response_payload, PROMPT_IDS_FIELD)

def extract_generation(self, response_payload: Any) -> tuple[list[int], list[float]]:
token_ids = _token_ids(response_payload, GENERATED_IDS_FIELD)
log_probs = _log_probs(response_payload, GENERATED_LOGPROBS_FIELD)
if len(token_ids) != len(log_probs):
raise ValueError(
f"Megatron generated token and log-probability lengths differ: {len(token_ids)} != {len(log_probs)}"
)
return token_ids, log_probs

def extract_extras(self, response_payload: Any) -> dict[str, Any] | None:
# Megatron Inference routed-experts rows are total_tokens - 1 long while the staging
# contract is delta-token aligned. Until that shift has a first-class
# representation the adapter stages no extras.
Comment thread
lauradang marked this conversation as resolved.
return None
7 changes: 6 additions & 1 deletion nemo_gym/token_id_capture/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -71,7 +71,7 @@
from collections.abc import Mapping
from importlib import import_module
from pathlib import Path
from typing import Any
from typing import Any, Literal

from pydantic import BaseModel, ConfigDict, Field, model_validator

Expand All @@ -86,6 +86,7 @@
logger = logging.getLogger(__name__)

TOKEN_ID_CAPTURE_BLOCK = "token_id_capture"
ExternalStagingBackend = Literal["vllm_worker", "megatron_worker"]


class TokenIdCaptureSettings(BaseModel):
Expand Down Expand Up @@ -133,6 +134,8 @@ class TokenIdCaptureSettings(BaseModel):
# It also makes committed parents visible to every serving worker.
# No additional in-memory coordinator is used.
external_staging: bool = False
# Both backends stage a canonical delta before returning coordinates.
external_staging_backend: ExternalStagingBackend = "vllm_worker"
# Name of the environment variable containing the manifest-route bearer token.
# The serving process reads the token without adding it to serialized configuration.
control_auth_token_env: str = Field(
Expand Down Expand Up @@ -170,6 +173,8 @@ def _validate(self) -> "TokenIdCaptureConfig":
"token_id_capture.external_staging requires rebuild_response=false because the "
"framework owns staged-record finalization"
)
if block.external_staging_backend == "megatron_worker" and not block.external_staging:
raise ValueError("token_id_capture.external_staging_backend requires external_staging=true")
if not block.enabled:
# Keep inactive settings for templated configurations.
# A run may toggle only ``enabled``.
Expand Down
Loading
Loading