Fréchet derivative of BatchNorm #
Spec.batchNorm flattens the spatial axes of a channel-first tensor to a [channels, positions]
matrix, normalizes every row of that matrix with the row statistics, and applies the per-channel
affine parameters. This file proves that, for 0 < ε, the map (x, gamma, beta) ↦ batchNorm is
Fréchet differentiable on the flattened vectors and that its derivative is exactly
Spec.batchNormJvp.
The calculus is delegated to RowNorm.hasFDerivAt_nrm, the derivative of one normalized row entry.
The rest of the file is bookkeeping: reshaping commutes with pointwise operations, the clamps in
the spec are inactive because the variance is a mean of squares, and every entry of the JVP is the
closed-form row differential.
Reshaping #
Reshaping commutes with pointwise binary operations.
Reshaping commutes with pointwise unary operations.
Reshaping a filled tensor fills the target shape.
Reshaping back and forth is the identity, for any proofs of the size equalities.
Reshaping commutes with tensor addition.
Reshaping commutes with tensor subtraction.
Reshaping commutes with pointwise multiplication.
Reshaping commutes with pointwise division.
Reshaping commutes with the clamped square root.
Transporting a vector along a length equality reindexes its Euclidean view.
tensorToVec of a reshaped tensor is a reindexing of tensorToVec.
Row means when the axis is written as rank - 1, as in the spec.
The flattened matrix form #
Number of positions per channel.
Instances For
Flattening the spatial axes preserves the size.
The number of positions is positive for a well-formed input shape.
The input flattened to a [channels, positions] matrix.
Instances For
Row broadcast of a channel vector to the flattened matrix shape.
Instances For
Spec.batchNorm with its axis and reshape evidence spelled out.
Instances For
Spec.batchNorm is batchNormExplicit by unfolding.
Entries of Spec.batchNorm in the flattened matrix: each row is centered, divided by the
clamped standard deviation, scaled and shifted by the channel parameters.
Spec.batchNormJvp with its axis and reshape evidence spelled out.
Instances For
Spec.batchNormJvp is batchNormJvpExplicit by unfolding.
Entries of the normalized-data JVP.
Vector form #
Domain of the vectorized BatchNorm: flattened input and the two channel parameters.
Instances For
Spec.batchNorm on flattened vectors.
Instances For
Flattened row index of a coordinate of the input.
Instances For
Flattened column index of a coordinate of the input.
Instances For
Channel index of a coordinate of the input, as an index into a channel vector.
Instances For
The input as a flattened channels × positions matrix vector.
Instances For
Every input coordinate is idxMN of its row and column, up to the size cast.
Entries of a flattened tensor are coordinates of matVec of its vectorization.
Entries of the flattened input tensor are coordinates of matVec.
Row means of a flattened tensor are RowNorm.rowMean of its vectorization.
Row variances of a flattened tensor are RowNorm.rowVar of its vectorization.
Entries of a channel vector are coordinates of its vectorization.
Closed form of the vectorized BatchNorm for positive ε.
Instances For
For positive ε the clamps in Spec.batchNorm are inactive and the vectorized map is its
closed form.
Graph.castCLM acts by castVec.
Derivative of the closed form at (X, g, _), as a continuous linear map.
Instances For
Coordinates of the batch-norm differential, split into its three contributions.
Reading the summands left to right: the normalization's own Jacobian applied to the input
perturbation and then scaled by γ, the normalized activation times the perturbation of γ, and
the perturbation of β. That decomposition is why the three parameter groups can be differentiated
independently downstream.
The closed form is differentiable for positive ε.
Main theorems #
Spec.batchNorm is Fréchet differentiable in (x, gamma, beta) for positive ε.
The vectorized BatchNorm at packed tensors is the vectorized spec output.
Entries of Spec.batchNormJvp in the flattened matrix are the closed-form row
differential, for positive ε.
The derivative of BatchNorm applied to packed tangents is Spec.batchNormJvp.
The Fréchet derivative of Spec.batchNorm in (x, gamma, beta) applied to a tangent triple
is Spec.batchNormJvp, for positive ε.