Skip to main content

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_gns every N steps (like 100+) to avoid overhead
  • Use an EMA on the GNS values since they are very noisy
  • The param_percentage param lets you sample a subset of model parameters for faster computation
  • Enable vmap with use_vmap=True to 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)

Source distribution for gns-pytorch 0.1.5
File Size Uploaded
gns_pytorch-0.1.5.tar.gz 29.4 kB Details

Built distribution (wheel)

Table of built distributions (wheels) for gns-pytorch 0.1.5
File Interpreter ABI Platform
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

Release history Release notifications | RSS feed

This release

0.1.5 This release

2 release files

0.1.4

2 release files

0.1.3

2 release files

0.1.2

2 release files

0.1.1

2 release files

0.1.0

2 release files

Anthropic, PBC Visionary sponsor Bloomberg Visionary sponsor Hudson River Trading Visionary sponsor Meta Visionary sponsor NVIDIA Visionary sponsor Microsoft Sustainability sponsor Depot Continuous Integration AWS Cloud computing and Security Sponsor Datadog Monitoring Fastly CDN Google Download Analytics Sentry Error logging StatusPage Status page