Download tests/test_normalization.py from Amitkumar001/Law_Slm: direct link, hf CLI and curl.
- Browser
- Download file 1.18 kB
-
https://huggingface.co/Amitkumar001/Law_Slm/resolve/main/tests/test_normalization.py
- Command line
-
hf download hf://Amitkumar001/Law_Slm/tests/test_normalization.py
-
curl -L -o test_normalization.py https://huggingface.co/Amitkumar001/Law_Slm/resolve/main/tests/test_normalization.py
1.18 kB
| """ | |
| Unit tests for RMSNorm and LayerNorm modules. | |
| """ | |
| import torch | |
| import pytest | |
| from slm.normalization.rmsnorm import RMSNorm | |
| from slm.normalization.layernorm import CustomLayerNorm | |
| def test_rmsnorm_forward_and_backward(): | |
| batch, seq, dim = 2, 8, 64 | |
| x = torch.randn(batch, seq, dim, requires_grad=True) | |
| norm = RMSNorm(dim=dim) | |
| out = norm(x) | |
| assert out.shape == (batch, seq, dim) | |
| # Check that root mean square across last dimension is close to 1.0 (before scaling) | |
| rms_val = torch.sqrt(out.pow(2).mean(dim=-1)) | |
| assert torch.allclose(rms_val, torch.ones_like(rms_val), atol=1e-2) | |
| loss = out.sum() | |
| loss.backward() | |
| assert x.grad is not None | |
| assert x.grad.shape == x.shape | |
| def test_layernorm_forward_and_backward(): | |
| batch, seq, dim = 2, 8, 64 | |
| x = torch.randn(batch, seq, dim, requires_grad=True) | |
| norm = CustomLayerNorm(dim=dim) | |
| out = norm(x) | |
| assert out.shape == (batch, seq, dim) | |
| # Check mean close to 0 and std close to 1 | |
| mean = out.mean(dim=-1) | |
| assert torch.allclose(mean, torch.zeros_like(mean), atol=1e-3) | |
| loss = out.sum() | |
| loss.backward() | |
| assert x.grad is not None | |