GNS PyTorch
This is the easiest way to calculate GNS (Gradient Noise Scale) for your PyTorch models. No hooks, gradient accumulation, or multi-GPU setup needed. Just pass in your per-example losses and model.
What's GNS?
GNS measures gradient noise in your training. See https://arxiv.org/pdf/1812.06162 and https://openreview.net/forum?id=xINTMAvPQA
Install
pip install gns-pytorch
Usage
Simple usage:
from gns_pytorch import compute_gns
import torch
model = YourModel()
optimizer = torch.optim.Adam(model.parameters())
def training_step(batch):
x, y = batch
logits = model(x)
per_example_losses = torch.nn.functional.cross_entropy(logits, y, reduction='none')
if global_step % 100 == 0:
gns_value = compute_gns(per_example_losses, model)
gns_ema = 0.9 * gns_ema + 0.1 * gns_value
print(f"Current GNS (EMA): {gns_ema}")
loss = per_example_losses.mean()
optimizer.zero_grad()
loss.backward()
optimizer.step()
Adaptive Batch Size Scheduling
With accurate GNS you can schedule your batch size (using gradient accumulation) to always be critical / optimal throughout training, massively boosting convergence and sample efficiency. This is similar to what deepseek-v3 did.
Tips
- Call
compute_gnsevery N steps (like 100+) to avoid overhead - Use an EMA on the GNS values since they are very noisy
- The
param_percentageparam lets you sample a subset of model parameters for faster computation - Enable vmap with
use_vmap=Trueto speed up computation by parallelizing per-example gradients (unfortunately, PyTorch's vmap isn't composable with flex attention and torch.compile yet) - GNS directly approximates critical batch size. For example, if GNS logger shows 64 and your global batch size is 32, you should double your gradient accumulation steps
Release files for gns-pytorch 0.1.5
For a detailed explanation of source distributions (sdists) and built distributions (wheels), please see the package formats documentation.
Source distribution (sdist)
| File | Size | Uploaded | |
|---|---|---|---|
| gns_pytorch-0.1.5.tar.gz | 29.4 kB | Details |
Built distribution (wheel)
| File | Interpreter | ABI | Platform | Reset |
|---|---|---|---|---|
| gns_pytorch-0.1.5-py3-none-any.whl | Python 3 | none | any | Details |
Total release size: 32.8 kB
Release files / gns_pytorch-0.1.5.tar.gz
| Download URL | gns_pytorch-0.1.5.tar.gz |
|---|---|
| Size | 29.4 kB |
| Tags | Source |
|
SHA-256 checksum How to use checksums |
2debe2e9966b1ea8b9b83dbd2d55a1a7c387e2f6ec55263a15afe869fa882b2c
|
|
BLAKE2b-256 checksum How to use checksums |
c61f77957a760e530e079922c3b907108175c1e68ce06e0cfd1737bab5b04529
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
No |
| Uploaded via |
uv/0.6.6
|
Release files / gns_pytorch-0.1.5-py3-none-any.whl
| Download URL | gns_pytorch-0.1.5-py3-none-any.whl |
|---|---|
| Size | 3.4 kB |
| Tags | Python 3 |
|
SHA-256 checksum How to use checksums |
9a6d2e054da8d75cb8e3bb97899c8991e8493008b4ed375dbd2f98f50608da79
|
|
BLAKE2b-256 checksum How to use checksums |
e03e2384236269a8d4183860d7d47453b71937052bc59dff0552b2437f7a3ff0
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
No |
| Uploaded via |
uv/0.6.6
|