Skip to content

Add dual-mode RoPE audio encoder - #7066

Open
desh2608 wants to merge 4 commits into
NVIDIA:mainfrom
desh2608:desh/gh-dual-mode-audio-encoder
Open

desh2608 wants to merge 4 commits into
NVIDIA:mainfrom
desh2608:desh/gh-dual-mode-audio-encoder

Conversation

@desh2608

@desh2608 desh2608 commented Sep 3, 2026 •

Copy link
Copy Markdown
  • I, the PR author, have personally reviewed every line of this PR.

What does this PR do?

Adds an opt-in dual-mode RoPE audio encoder that uses one shared encoder parameterization for offline and causal audio processing.

The implementation:

  • adds a RoPE Transformer audio encoder with frame stacking, QK normalization, pre-block normalization, checkpointing, and bounded causal left context;
  • supports dense and packed THD inputs;
  • supports offline, causal, and per-sequence mixed-mode attention within the same packed batch without reordering tokens;
  • constructs the mixed THD attention policy once and reuses it across encoder layers;
  • integrates the encoder with the existing NeMo audio model and checkpoint/config conversion paths;
  • unifies legacy and RoPE frame stacking behind one pre_encode="stacking" implementation with configurable learned or zero padding;
  • preserves the existing relative-position audio encoder as the default and requires explicit architecture=rope_transformer selection for the new path.

TransformerEngine dependency

Dual-mode training requires the latest TransformerEngine with the grouped mixed-THD attention-policy API from NVIDIA/TransformerEngine#3274. The mixed packed path passes thd_attention_policies with thd_attention_policy_dispatch="grouped"; TransformerEngine versions that predate that API are not sufficient.

Configuration and compatibility

  • NeMo archives with self_attention_model: rope are recognized as rope_transformer checkpoints.
  • NeMo feat_in maps to the encoder n_mels setting.
  • Both encoder architectures now use pre_encode="stacking" when frame stacking is selected.
  • stacking_pad_mode="learned" preserves the legacy NGPT padding behavior; stacking_pad_mode="zeros" matches the RoPE NeMo encoder.
  • The former feature_stacking value remains accepted as a compatibility alias and normalizes to zero-padded stacking.
  • Existing state-dict layouts remain unchanged: learned padding uses proj_out.weight plus pad_frame, while zero padding uses proj.weight without a trainable padding parameter.
  • --nemo-transformer-audio-causal-mode supports checkpoint, causal, and offline runtime selection.
  • Legacy relative-position checkpoints retain their existing static causal/offline behavior.

Issue tracking

Linked issue: N/A — internal feature port.

Validation

Focused unit tests:

python -m pytest -q \
  tests/unit_tests/models/test_nemo_rope_transformer_audio.py \
  tests/unit_tests/models/test_nemo_transformer_audio.py

Result: 40 passed, 6 skipped. The skipped cases require CUDA and Transformer Engine; the suite includes a real-TE CUDA parity test for mixed windowed packed attention.

The stacking tests cover learned and zero padding on dense and already-packed inputs, padding numerics, checkpoint-key compatibility, legacy config normalization, and RoPE configuration validation.

Additional checks completed:

  • tools/autoformat.sh in check mode
  • isort and Black formatting
  • Pylint: 10.00/10
  • Ruff: all checks passed
  • Python bytecode compilation
  • git diff --check

Contribution process

Pre-checks

  • I have added relevant unit tests
  • I have added relevant functional tests — the behavior is covered by focused CPU tests and a CUDA/Transformer Engine parity unit test
  • I have added proper typing to my code Typing guidelines
  • I have added relevant documentation through configuration and API docstrings
  • I have run tools/autoformat.sh on my PR

This PR changes attention/model architecture behavior and adds new tests, so it uses the Run functional tests label.

@copy-pr-bot

copy-pr-bot Bot commented Sep 3, 2026

Copy link
Copy Markdown

This pull request requires additional validation before any workflows can run on NVIDIA's runners.

Pull request vetters can view their responsibilities here.

Contributors can view more details about this message here.

@desh2608
desh2608 marked this pull request as ready for review September 3, 2026 16:30
@desh2608
desh2608 requested a review from a team as a code owner September 3, 2026 16:30

@yqwangustc yqwangustc left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Changes look good to me. Only some minor comment on docstring for the "feature_stacking" vs "stacking", and "nemo_transformer_audio_causal_mode==checkpoint" means.

If "feature_stacking" is very similar to "stacking", and the only difference is how they treat padding when doing the stacking, maybe consider merge them to one mode ?

return num_frames // self.encoder_time_stride

if self.pre_encode == "stacking":
if self.pre_encode in ("stacking", "feature_stacking"):

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Do we have any docstring to explain what's the difference between "stacking" and "feature_stacking" ?

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Removed "feature_stacking" in 15fe07e. Now both legacy and rope encoders use "stacking" mode, just with different stacking_pad_mode.


raise ValueError(
f"Unsupported Nemo TransformerEncoder pre_encode={self.pre_encode!r}; "
"expected 'conv', 'depth_conv', or 'stacking'."

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

looks like here should be "expected 'conv', 'depth_conv', 'stacking' or 'feature_stacking' ?

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Removed.

left_context = getattr(args, "nemo_transformer_audio_left_context", None)
if left_context is not None:
encoder_cfg.left_context = None if left_context < 0 else left_context
causal_mode = getattr(args, "nemo_transformer_audio_causal_mode", "checkpoint")

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Do we have a docstring to explain what "nemo_transformer_audio_causal_mode=checkpoint" mean ?

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Added in 2add54b

@svcnvidia-nemo-ci svcnvidia-nemo-ci added the Approved All necessary approvals have been made label Sep 3, 2026
@desh2608
desh2608 enabled auto-merge September 3, 2026 21:50
@desh2608

desh2608 commented Sep 4, 2026

Copy link
Copy Markdown
Author

@yqwangustc not sure why some checks are indefinitely pending, do you know?

@yqwangustc

Copy link
Copy Markdown
Contributor

/ok to test 8100f03

Comment thread pyproject.toml Outdated
[[tool.uv.dependency-metadata]]
name = "transformer-engine"
version = "2.18.0+27486e03"
version = "2.20.0.dev0+7ba2f9f1"

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The new TE version is required which enables thd_attention_policies. See NVIDIA/TransformerEngine#3274

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Looks like this change has landed. Are we safe to make this pyproject.toml file changes now ?

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yeah it should be safe, but we need code-owner approval.

@yqwangustc
yqwangustc requested a review from a team September 24, 2026 03:56
@desh2608
desh2608 force-pushed the desh/gh-dual-mode-audio-encoder branch from 7c989cf to cd3c099 Compare September 24, 2026 14:01
@yqwangustc

Copy link
Copy Markdown
Contributor

/ok to test cd3c099

Port the dual-mode causal/offline audio encoder onto the current GitHub audio model interfaces. Preserve one encoder parameterization across modes and use Transformer Engine's grouped THD attention-policy dispatch for mixed packed batches.

Signed-off-by: Desh Raj <deshr@nvidia.com>
Use one stacking pre-encoder for legacy and RoPE audio towers, with an explicit learned-or-zero padding policy. Preserve both checkpoint layouts and normalize the former feature_stacking configuration as a compatibility alias.

Signed-off-by: Desh Raj <deshr@nvidia.com>
Signed-off-by: Desh Raj <deshr@nvidia.com>
@desh2608
desh2608 force-pushed the desh/gh-dual-mode-audio-encoder branch from cd3c099 to e22de2a Compare October 1, 2026 17:24
@copy-pr-bot

copy-pr-bot Bot commented Oct 1, 2026

Copy link
Copy Markdown

/ok to test

@desh2608, there was an error processing your request: E1

See the following link for more information: https://docs.gha-runners.nvidia.com/cpr/e/1/

@desh2608

desh2608 commented Oct 1, 2026

Copy link
Copy Markdown
Author

/ok to test e22de2a

@desh2608

desh2608 commented Oct 2, 2026

Copy link
Copy Markdown
Author

/ok to test 0af3991

@desh2608
desh2608 added this pull request to the merge queue Oct 2, 2026
@nemo-automation-bot

Copy link
Copy Markdown

🔄 Merge queue validation started!

You can track the progress here: https://github.com/NVIDIA/Megatron-LM/actions/runs/37052597572

@github-merge-queue
github-merge-queue Bot removed this pull request from the merge queue due to failed status checks Oct 2, 2026

This branch was successfully deployed

2 active deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

Approved All necessary approvals have been made

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants