Skip to content

Use native RMSNorm directly in REVE and ZUNA - #1176

Merged
bruAristimunha merged 1 commit into
masterfrom
simplify-native-rmsnorm
Sep 21, 2026
Merged

bruAristimunha merged 1 commit into
masterfrom
simplify-native-rmsnorm

Conversation

@bruAristimunha

Copy link
Copy Markdown
Collaborator

REVE and ZUNA now import torch.nn.RMSNorm directly. Remove both local subclasses and their adapter-specific tests now that #1174 has raised the minimum PyTorch version to 2.4. Keep each model's explicit epsilon and checkpoint parameter names.

Validation: all configured pre-commit hooks passed; PyTorch 2.4.0 ran four focused attention/rotary tests and 11 model integration tests (14 existing skips). Small models strictly load the previous checkpoints and match float32 logits and gradients exactly. CPU bfloat16 forward/backward remains finite; ZUNA logits differed by at most 0.0009765625 in the seeded comparison.

Mixed precision now uses native PyTorch semantics. The removed adapters' float32 accumulation guarantee no longer applies; PyTorch 2.4's native float16 RMSNorm can overflow for large inputs.

Copilot AI lite review requested due to automatic review settings September 21, 2026 19:10
@bruAristimunha
bruAristimunha merged commit 4e332c4 into master Sep 21, 2026
11 of 12 checks passed

Copilot AI 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.

Copilot review overview

🟡 Changes recommended

The removal of RMSNorm adapter tests leaves a gap in coverage for the new invariants (native type/eps/checkpoint keys), and the updated release-note entry appears to reference the wrong PR number.

Get a fresh assessment by requesting another Copilot review.

Review effort: Lite
Findings: 1 Medium severity · 1 Low severity

Open (2)
What changed in this PR

This pull request updates the REVE and ZUNA foundation models to rely on PyTorch’s native torch.nn.RMSNorm directly (now that PyTorch >= 2.4 is required), removing the project’s adapter/subclass implementations and the associated precision-focused unit test.

Changes:

  • Swap REVE and ZUNA internal normalization layers to use torch.nn.RMSNorm directly while keeping explicit epsilon values.
  • Remove the adapter-specific RMSNorm precision/gradient reference test.
  • Update the Requirements release-note wording to reflect the native RMSNorm usage.
File Description
braindecode/​models/​reve.py Replace local RMSNorm subclass usage with direct import/use of torch.nn.RMSNorm and keep eps=1e-6 explicit.
braindecode/​models/​zuna.py Remove custom _RMSNorm adapter and switch all norm sites to torch.nn.RMSNorm.
test/​unit_tests/​models/​test_foundation_models.py Drop adapter-specific RMSNorm precision test and remove now-unused import.
docs/​whats_new.rst Adjust Requirements entry text to describe native RMSNorm import and explicit eps behavior.

💡 Add a code-review agent skill or configure MCP servers for context-aware, tailored reviews. Learn more in the docs.

Comment on lines 726 to 728


@pytest.mark.parametrize(
"norm_class, eps", [(RMSNorm, 1e-6), (zuna_module._RMSNorm, 1e-5)]
)
@pytest.mark.parametrize("input_dtype", [torch.float32, torch.float16, torch.bfloat16])
@pytest.mark.parametrize("weight_dtype", [torch.float32, torch.float16, torch.bfloat16])
def test_foundation_rms_norm_preserves_reference_precision(
norm_class, eps, input_dtype, weight_dtype
):
norm = RMSNorm(dim=8) if norm_class is RMSNorm else norm_class(8, eps=eps)
norm = norm.to(weight_dtype)
weight = torch.linspace(0.5, 1.5, 8, dtype=weight_dtype, requires_grad=True)
norm.load_state_dict({"weight": weight.detach()}, strict=True)
# Squaring 1000 overflows float16; small values exercise the explicit eps.
x = torch.linspace(-1, 1, 24).reshape(3, 8)
x = (
(x * torch.tensor([1e-4, 1.0, 1000.0])[:, None])
.to(input_dtype)
.requires_grad_()
)
reference_x = x.detach().clone().requires_grad_()
normalized = reference_x.float() * torch.rsqrt(
reference_x.float().square().mean(-1, keepdim=True) + eps
)
dtype = input_dtype if norm_class is RMSNorm else weight_dtype
expected = normalized.to(dtype) * weight
actual = norm(x)
torch.testing.assert_close(actual, expected)
actual.sum().backward()
expected.sum().backward()
torch.testing.assert_close(x.grad, reference_x.grad)
torch.testing.assert_close(norm.weight.grad, weight.grad)


def test_reve_attention_matches_explicit_attention():
Comment thread docs/whats_new.rst
Comment on lines 43 to 47
- Require PyTorch and TorchAudio >= 2.4 and remove obsolete attention fallbacks.
RMS normalization in REVE and ZUNA now uses PyTorch's implementation while
retaining float32 accumulation. Intel macOS is no longer supported because
REVE and ZUNA now import PyTorch's RMSNorm layer directly, preserving their
explicit epsilon values. Intel macOS is no longer supported because
PyTorch stopped providing its binary packages after 2.2.
(:gh:`1174` by `Bruno Aristimunha`_)
@codecov

codecov Bot commented Sep 21, 2026

Copy link
Copy Markdown

Codecov Report

✅ All modified and coverable lines are covered by tests.
✅ Project coverage is 86.78%. Comparing base (25c7ed6) to head (ec7d18a).
⚠️ Report is 2 commits behind head on master.

Additional details and impacted files
@@            Coverage Diff             @@
##           master    #1176      +/-   ##
==========================================
+ Coverage   86.57%   86.78%   +0.21%     
==========================================
  Files         142      142              
  Lines       16303    16262      -41     
==========================================
- Hits        14114    14113       -1     
+ Misses       2189     2149      -40     
🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.

bruAristimunha added a commit to julien-gadonneix/braindecode that referenced this pull request Sep 24, 2026
Take master's REVE and ZUNA files (same RMSNorm change, from braindecode#1176) and
drop the duplicate API-changes entry for the PyTorch >= 2.4 requirement
(braindecode#1174). Affected model tests: 328 passed, 56 skipped.
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants