-
Notifications
You must be signed in to change notification settings - Fork 360
feat(token-capture): add Megatron worker capture backend #2823
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Merged
ananthsub
merged 15 commits into
NVIDIA-NeMo:main
from
lauradang:laurad/megatron-adapter
Sep 19, 2026
Merged
Changes from all commits
Commits
Show all changes
15 commits
Select commit
Hold shift + click to select a range
5f4ddab
feat(token-capture): add Megatron worker capture backend
lauradang fc08bf1
Merge branch 'main' into laurad/megatron-adapter
lauradang a1897b6
Merge branch 'main' into laurad/megatron-adapter
lauradang 7e89d2f
feat(token-capture): add Gym-owned Megatron extraction adapter
lauradang ba47b88
refactor(token-capture): share acknowledgement commit path across bac…
lauradang 5f46ece
fix(token-capture): share the invalid-coordinates failure reason acro…
lauradang 91250ec
test(token-capture): exercise the acknowledgement contract on both ha…
lauradang a3cef77
fix(token-capture): send offload_params to the Megatron endpoint
lauradang 37dc751
fix(token-capture): allow RolloutReceipt.terminal_selection to be unset
lauradang 04f0354
refactor(token-capture): stop requesting return_tokenized_data from M…
lauradang d8207bd
Merge branch 'main' into laurad/megatron-adapter
lauradang 8293c3f
Update nemo_gym/token_id_capture/adapters/megatron.py
lauradang c04f7ef
fix(token-capture): enforce RolloutReceipt attribution-state invariants
lauradang f11130e
fix(token-capture): fail closed on malformed or multimodal Megatron c…
lauradang b601670
refactor(token-capture): name the backend in the external_staging con…
lauradang File filter
Filter by extension
Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
There are no files selected for viewing
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| 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. | ||
| return None | ||
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Oops, something went wrong.
Oops, something went wrong.
Add this suggestion to a batch that can be applied as a single commit.
This suggestion is invalid because no changes were made to the code.
Suggestions cannot be applied while the pull request is closed.
Suggestions cannot be applied while viewing a subset of changes.
Only one suggestion per line can be applied in a batch.
Add this suggestion to a batch that can be applied as a single commit.
Applying suggestions on deleted lines is not supported.
You must change the existing code in this line in order to create a valid suggestion.
Outdated suggestions cannot be applied.
This suggestion has been applied or marked resolved.
Suggestions cannot be applied from pending reviews.
Suggestions cannot be applied on multi-line comments.
Suggestions cannot be applied while the pull request is queued to merge.
Suggestion cannot be applied right now. Please check back later.
Uh oh!
There was an error while loading. Please reload this page.