Repository navigation
Use native RMSNorm directly in REVE and ZUNA - #1176
Conversation
There was a problem hiding this comment.
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
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.RMSNormdirectly 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.
|
|
||
|
|
||
| @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(): |
| - 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 Report✅ All modified and coverable lines are covered by tests. 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:
|
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.


REVE and ZUNA now import
torch.nn.RMSNormdirectly. 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.