NVIDIA / NVIDIA/Megatron-LM

Reduce downstream Megatron patching for RL use cases

Open
#4,590 15 comments 1 reaction 0 assignees View on GitHub
enhancement module: rl
Dominant language
Python
Stars
17.9k
Forks
4.5k
Avg merge
4d 6h
Merged PRs (30d)
271

Description

Some Megatron Core features are difficult to use from external RL training loops without copying or monkey-patching GPTModel.forward, GPT postprocess, MTP postprocess, or 1F1B schedule plan.

This usually happens when the training loop owns data or semantics that Megatron Core should not model directly: selected-token labels, loss masks, packed sequence metadata, old/reference logprobs, KL/entropy terms, or custom fused logprob/loss computation.

**Current downstream symptoms**

- veRL patches/copies GPT/MTP postprocess logic:
- verl/models/mcore/mtp_patch.py
- verl/models/mcore/model_forward_fused.py
- verl/models/mcore/model_forward_1f1b_overlap.py

This indicates that external training loops need a stable, objective-neutral extension point at the GPT postprocess boundary.

**Proposed direction**

Add a small optional GPT output/postprocess hook, keyword-only. The hook should run after decoder hidden states are available and before the default output-layer logits/loss path. This should avoid adding PPO/GRPO/RL-specific arguments to GPTModel.forward.

Schedule-plan support

Thread the same optional processor/context through build_schedule_plan and the 1f1b schedule-plan PostProcessNode.

MTP follow-up

Handle MTP separately if needed. First investigate whether MTP can expose a narrow callable for custom loss/logprob computation while Megatron Core continues to own MTP shifting, packed-sequence handling, scaling, and logging behavior.

Contributor guide

Open the contributing guide

Assessment

This issue has not been assessed yet.

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.