[megatron] fix: forward mhc_multistream to MTP and skip activation reclaim for MTP checkpoints - #7328
Merged
Conversation
…claim for MTP checkpoints Fixes two crashes that block mHC (manifold hyper-connections) + MTP training on Megatron-core (DeepSeek-V4, `use_fused_mhc=False` + `mtp.enable=True`). Both stem from verl's MTP postprocess / recomputation patches not accounting for the mHC multi-stream tensor. **Background:** With mHC + MTP, the decoder returns a tuple `(contracted_hidden, mhc_multistream)` where `mhc_multistream` is the pre-contraction `[s, b, n*h]` multi-stream tensor. `GPTModel._postprocess` forwards `mhc_multistream` into `self.mtp(...)`, and the MTP module branches on it (`if mhc_multistream is not None`) to use the multi-stream tensor as its depth input. `mhc_multistream` is a plain alias of the decoder output (`transformer_block.py`: `mhc_multistream = hidden_states`), and the MTP depth input is a `torch.chunk` **view** of it — i.e. shared storage, not a copy. **Bug 1 — reshape crash (forward):** verl's patched `_megatron_gptmodel_postprocess` (`mtp_patch.py`) accepted `mhc_multistream` as a parameter but never forwarded it into `self.mtp(...)`, so it arrived as `None` inside MTP. MTP then took the contracted `[s, b, h]` path while the mHC branch of `_concat_embeddings` expected `[s, b, n*h]`, failing at: ``` RuntimeError: shape '[23986, 1, 4, 4096]' is invalid for input of size 98246656 ``` Fix: forward `mhc_multistream` into the `self.mtp(...)` call, matching the upstream `GPTModel._postprocess`. **Bug 2 — async CUDA illegal-memory-access (backward):** verl's recomputation backward patch (`apply_patch_megatron_recomputation_backward`, commit 04df110, a MoE residual-memory leak fix) reclaims saved activation memory by calling `t.untyped_storage().resize_(0)` on every checkpoint input after computing grads. With mHC + MTP this is a **use-after-free**: the MTP checkpoint's saved `hidden_states` is a `torch.chunk` view of the decoder's `mhc_multistream` (shared storage via `make_viewless_tensor` → `_kernel_make_viewless_tensor` doing `out.data = inp.data`), and the MTP checkpoint backward runs **before** the decoder's `learned_output_contract` backward (reverse topological order). `resize_(0)` therefore truncates storage the pending decoder backward still reads → async CUDA illegal-memory-access, surfacing at the next sync point (`set_state` inside `_fork_rng.__exit__`). Fix: detect MTP-layer checkpoints by inspecting the `MultiTokenPredictionLayer` instance captured in `ctx.run_function.__closure__` (its `__qualname__` is `..._checkpointed_forward.<locals>.custom_forward`, which carries no class name), and skip the `resize_(0)` reclaim for those. Non-MTP checkpoints keep the reclaim, preserving the original MoE leak fix. Signed-off-by: Hollow Man <hollowman@opensuse.org>
Signed-off-by: Hollow Man <hollowman@opensuse.org>
wuxibin89
approved these changes
Aug 10, 2026
kahlun
pushed a commit
to kahlun/verl
that referenced
this pull request
Aug 20, 2026
…claim for MTP checkpoints (verl-project#7328) ### What does this PR do? Fixes two crashes that block mHC (manifold hyper-connections) + MTP training on Megatron-core (DeepSeek-V4, `use_fused_mhc=False` + `mtp.enable=True`). Both stem from verl's MTP postprocess / recomputation patches not accounting for the mHC multi-stream tensor. **Background:** With mHC + MTP, the decoder returns a tuple `(contracted_hidden, mhc_multistream)` where `mhc_multistream` is the pre-contraction `[s, b, n*h]` multi-stream tensor. `GPTModel._postprocess` forwards `mhc_multistream` into `self.mtp(...)`, and the MTP module branches on it (`if mhc_multistream is not None`) to use the multi-stream tensor as its depth input. `mhc_multistream` is a plain alias of the decoder output (`transformer_block.py`: `mhc_multistream = hidden_states`), and the MTP depth input is a `torch.chunk` **view** of it — i.e. shared storage, not a copy. **Bug 1 — reshape crash (forward):** verl's patched `_megatron_gptmodel_postprocess` (`mtp_patch.py`) accepted `mhc_multistream` as a parameter but never forwarded it into `self.mtp(...)`, so it arrived as `None` inside MTP. MTP then took the contracted `[s, b, h]` path while the mHC branch of `_concat_embeddings` expected `[s, b, n*h]`, failing at: ``` RuntimeError: shape '[23986, 1, 4, 4096]' is invalid for input of size 98246656 ``` Fix: forward `mhc_multistream` into the `self.mtp(...)` call, matching the upstream `GPTModel._postprocess`. **Bug 2 — async CUDA illegal-memory-access (backward):** verl's recomputation backward patch (`apply_patch_megatron_recomputation_backward`, commit 04df110, a MoE residual-memory leak fix) reclaims saved activation memory by calling `t.untyped_storage().resize_(0)` on every checkpoint input after computing grads. With mHC + MTP this is a **use-after-free**: the MTP checkpoint's saved `hidden_states` is a `torch.chunk` view of the decoder's `mhc_multistream` (shared storage via `make_viewless_tensor` → `_kernel_make_viewless_tensor` doing `out.data = inp.data`), and the MTP checkpoint backward runs **before** the decoder's `learned_output_contract` backward (reverse topological order). `resize_(0)` therefore truncates storage the pending decoder backward still reads → async CUDA illegal-memory-access, surfacing at the next sync point (`set_state` inside `_fork_rng.__exit__`). Fix: detect MTP-layer checkpoints by inspecting the `MultiTokenPredictionLayer` instance captured in `ctx.run_function.__closure__` (its `__qualname__` is `..._checkpointed_forward.<locals>.custom_forward`, which carries no class name), and skip the `resize_(0)` reclaim for those. Non-MTP checkpoints keep the reclaim, preserving the original MoE leak fix. ### Checklist Before Starting - [X] Search for similar PRs. Paste at least one query link here: ... - [X] Format the PR title as `[{modules}] {type}: {description}` (This will be checked by the CI) - `{modules}` include `fsdp`, `megatron`, `veomni`, `sglang`, `vllm`, `rollout`, `trainer`, `ci`, `training_utils`, `recipe`, `hardware`, `deployment`, `ray`, `worker`, `single_controller`, `misc`, `perf`, `model`, `algo`, `env`, `tool`, `ckpt`, `doc`, `data`, `cfg`, `reward`, `fully_async`, `one_step_off` - If this PR involves multiple modules, separate them with `,` like `[megatron, fsdp, doc]` - `{type}` is in `feat`, `fix`, `refactor`, `chore`, `test` - If this PR breaks any API (CLI arguments, config, function signature, etc.), add `[BREAKING]` to the beginning of the title. - Example: `[BREAKING][fsdp, megatron] feat: dynamic batching` ### Test > For changes that can not be tested by CI (e.g., algorithm implementation, new model support), validate by experiment(s) and show results like training curve plots, evaluation results, etc. ### API and Usage Example > Demonstrate how the API changes if any, and provide usage example(s) if possible. ```python # Add code snippet or script demonstrating how to use this ``` ### Design & Code Changes > Demonstrate the high-level design if this PR is complex, and list the specific changes. ### Checklist Before Submitting > [!IMPORTANT] > Please check all the following items before requesting a review, otherwise the reviewer might deprioritize this PR for review. - [X] Read the [Contribute Guide](https://github.com/verl-project/verl/blob/main/CONTRIBUTING.md). - [X] Apply [pre-commit checks](https://github.com/verl-project/verl/blob/main/CONTRIBUTING.md#code-linting-and-formatting): `pre-commit install && pre-commit run --all-files --show-diff-on-failure --color=always` - [X] Add / Update [the documentation](https://github.com/verl-project/verl/tree/main/docs). - [X] Add unit or end-to-end test(s) to [the CI workflow](https://github.com/verl-project/verl/tree/main/.github/workflows) to cover all the code. If not feasible, explain why: ... - [X] Once your PR is ready for CI, send a message in [the `ci-request` channel](https://verl-project.slack.com/archives/C091TCESWB1) in [the `verl` Slack workspace](https://join.slack.com/t/verl-project/shared_invite/zt-3855yhg8g-CTkqXu~hKojPCmo7k_yXTQ). (If not accessible, please try [the Feishu group (飞书群)](https://applink.larkoffice.com/client/chat/chatter/add_by_link?link_token=772jd4f1-cd91-441e-a820-498c6614126a).) - [X] If your PR is related to the `recipe` submodule, please also update the reference to the submodule commit via `git submodule update --remote` or `cd recipe && git pull origin main`. --------- Signed-off-by: Hollow Man <hollowman@opensuse.org>
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
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
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.
What does this PR do?
Fixes two crashes that block mHC (manifold hyper-connections) + MTP training on Megatron-core (DeepSeek-V4,
use_fused_mhc=False+mtp.enable=True). Both stem from verl's MTP postprocess / recomputation patches not accounting for the mHC multi-stream tensor.Background: With mHC + MTP, the decoder returns a tuple
(contracted_hidden, mhc_multistream)wheremhc_multistreamis the pre-contraction[s, b, n*h]multi-stream tensor.GPTModel._postprocessforwardsmhc_multistreamintoself.mtp(...), and the MTP module branches on it (if mhc_multistream is not None) to use the multi-stream tensor as its depth input.mhc_multistreamis a plain alias of the decoder output (transformer_block.py:mhc_multistream = hidden_states), and the MTP depth input is atorch.chunkview of it — i.e. shared storage, not a copy.Bug 1 — reshape crash (forward): verl's patched
_megatron_gptmodel_postprocess(mtp_patch.py) acceptedmhc_multistreamas a parameter but never forwarded it intoself.mtp(...), so it arrived asNoneinside MTP. MTP then took the contracted[s, b, h]path while the mHC branch of_concat_embeddingsexpected[s, b, n*h], failing at:Fix: forward
mhc_multistreaminto theself.mtp(...)call, matching the upstreamGPTModel._postprocess.Bug 2 — async CUDA illegal-memory-access (backward): verl's recomputation backward patch (
apply_patch_megatron_recomputation_backward, commit 04df110, a MoE residual-memory leak fix) reclaims saved activation memory by callingt.untyped_storage().resize_(0)on every checkpoint input after computing grads. With mHC + MTP this is a use-after-free: the MTP checkpoint's savedhidden_statesis atorch.chunkview of the decoder'smhc_multistream(shared storage viamake_viewless_tensor→_kernel_make_viewless_tensordoingout.data = inp.data), and the MTP checkpoint backward runs before the decoder'slearned_output_contractbackward (reverse topological order).resize_(0)therefore truncates storage the pending decoder backward still reads → async CUDA illegal-memory-access, surfacing at the next sync point (set_stateinside_fork_rng.__exit__).Fix: detect MTP-layer checkpoints by inspecting the
MultiTokenPredictionLayerinstance captured inctx.run_function.__closure__(its__qualname__is..._checkpointed_forward.<locals>.custom_forward, which carries no class name), and skip theresize_(0)reclaim for those. Non-MTP checkpoints keep the reclaim, preserving the original MoE leak fix.Checklist Before Starting
[{modules}] {type}: {description}(This will be checked by the CI){modules}includefsdp,megatron,veomni,sglang,vllm,rollout,trainer,ci,training_utils,recipe,hardware,deployment,ray,worker,single_controller,misc,perf,model,algo,env,tool,ckpt,doc,data,cfg,reward,fully_async,one_step_off,like[megatron, fsdp, doc]{type}is infeat,fix,refactor,chore,test[BREAKING]to the beginning of the title.[BREAKING][fsdp, megatron] feat: dynamic batchingTest
API and Usage Example
# Add code snippet or script demonstrating how to use thisDesign & Code Changes
Checklist Before Submitting
Important
Please check all the following items before requesting a review, otherwise the reviewer might deprioritize this PR for review.
pre-commit install && pre-commit run --all-files --show-diff-on-failure --color=alwaysci-requestchannel in theverlSlack workspace. (If not accessible, please try the Feishu group (飞书群).)recipesubmodule, please also update the reference to the submodule commit viagit submodule update --remoteorcd recipe && git pull origin main.