latest (dev)
Copy
Latest development documentation · Updated 2026-10-08
Patching batch norm
What’s happening?
Batch normalization updates its running_mean and running_var buffers in
place on every training forward. Under a transform that evaluates a module
more than once — differentiating through it, or mapping an ensemble over it
— those updates would be applied repeatedly and in an order the caller
never asked for. The running statistics are also not inputs of the
function, so no transform can account for them.
What works today
The operator-coverage picture for normalization, verified on the current build:
vmap()over a batch norm module raisesNotImplementedError(Kernel not found for op: batch_norm on backend: VmapCPU) — in training mode, in eval mode, and after the patching described below.group_normandlayer_normare in the same state, so switching the norm layer does not help undervmap()yet.grad()through a batch norm module does work: the forward is evaluated once, the batch statistics differentiate normally. The running-stat buffers are mutated as a side effect — once per evaluation of the transformed function.jacrev()through a batch norm module also works: the per-sample Jacobians are computed from the batch statistics of each evaluation and come out finite and correct. Because the running-stat buffers mutate on each pass, the values you get depend on how many times the transformed function ran before the call — use the helpers below when the model doubles as an inference network.
Warning
The Jacobian transforms evaluate the module once per output entry, so a batch norm’s running statistics move between evaluations. That drift does not corrupt the transform itself — the batch statistics of each call are what get differentiated — but it does mean the module’s buffers are unpredictable afterwards. Patch the norm out first if the model must stay state-free.
The helpers
replace_all_batch_norm_modules_() drops the running statistics of
every batch-normalization module in a tree, in place, and returns the
root. After the call, track_running_stats is False and the
running_mean / running_var / num_batches_tracked buffers are None:
the module normalizes with the batch’s own statistics and no longer
mutates state on the forward path.
import tensorplay as tp
import tensorplay.nn as nn
from tensorplay.func import grad, replace_all_batch_norm_modules_
net = nn.Sequential(nn.Conv2d(3, 4, 3), nn.BatchNorm2d(4))
replace_all_batch_norm_modules_(net)
assert net[1].track_running_stats is False
assert net[1].running_mean is None
x = tp.randn(2, 3, 8, 8)
g = grad(lambda t: net(t).sum())(x) # mutation-free evaluation
assert g.shape == x.shape
batch_norm_without_running_stats() is the single-module form: it
drops the running statistics of one batch-normalization module, if it is
one.
When to use them
Under the single-pass transforms (
grad(),grad_and_value(),vjp()): patching makes the evaluation free of hidden state — the same input always produces the same output, which is what you want when differentiating.Under
vmap(): normalization layers are not transformable today regardless of the running-stats setting. Strip or replace them before mapping; the helpers keep the module ready for when the rules land.Under the Jacobian transforms:
jacrev()works but leaves the running-stat buffers mid-stream, andjacfwd()is still covered by the general limitation above — patch the norm out to keep the transformed model state-free.
Help improve this page
Found an error, an unclear step, or a missing example?
Multiprocessing package - tensorplay.multiprocessing
Warning
Pruning - tensorplay.ao.pruning
tensorplay.ao.pruning sparsifies a trained model: it zeroes weights that contribute least, so the model stays the same size but computes less (with sparse-aware kernels) or compresses better on disk. Pruning is implement

