ParamGradNodes — each batched parameter gradient node is a loss derivative #
GradNodesB states each emitted parameter gradient node as its layer's parameter Jacobian
contracted with an arbitrary output cotangent. Here the cotangent is the gradient of a scalar G
at the layer's output (HasGradAt), and the node becomes ∂G/∂θ with the layer's parameter
varied: one lemma per node kind, shared by every net.
The BatchNorm γ/β nodes are stated in the transposed [C, N·H·W] layout at the reassocB
index; their lemmas re-sum the Jacobian over that permutation (bnLA_perm). A conv or depthwise
bias that the render reads with the β op (EfficientNet-B0's) is biasBeta_eq_pdiv: the bias enters
as a channel broadcast (*_bias_split), so the channel sum the β node computes is its derivative.
cStridedInB — the emitted strided-conv input-cotangent — is the batched strided conv VJP's
backward, at any saved input.
dInB — the emitted depthwise input-cotangent — is the batched depthwise VJP's backward.
dStridedInB — the emitted (symmetric) strided depthwise input-cotangent — is the batched
strided depthwise VJP's backward, at any saved input.
Back through batch BN: the gradient at its input is bnInB of the gradient at its output.
Back through batch BN, at the certified backward's own spelling bnBackB (EfficientNet-B0's
tie threads this form; the ResNets and MobileNets emit bnInB).
Back through a batched conv: cInB.
Back through a batched symmetric strided conv: cStridedInB.
Back through a batched depthwise: dInB.
Back through a batched symmetric strided depthwise: dStridedInB.
Conv weight node = ∂G/∂W.
Conv bias node = ∂G/∂b.
Stride-2 (symmetric) conv weight node = ∂G/∂W.
Stride-2 (symmetric) conv bias node = ∂G/∂b.
XLA-SAME strided conv weight node = ∂G/∂W (MobileNetV2's and B0's stem).
XLA-SAME strided conv bias node = ∂G/∂b.
Depthwise weight node = ∂G/∂W.
Depthwise bias node = ∂G/∂b.
Symmetric strided depthwise weight node = ∂G/∂W (MobileNetV4's dw_mid at its three
downsampling rows).
XLA-SAME strided depthwise weight node = ∂G/∂W.
XLA-SAME strided depthwise bias node = ∂G/∂b.
Dense weight node = ∂G/∂W.
Dense bias node = ∂G/∂b. The node's statement carries one activation x₀ for every
example; the bias Jacobian is the identity whatever the activation, so any x₀ serves.
The permutation bnBatchLA reads its per-channel core through: network index J ↦ the
[C, N·H·W] cell bnchwBackIdx (J at the mul_assoc cast).
Equations
- One or more equations did not get rendered due to their size.
Instances For
bnBatchLA at a network index IS the per-channel core at the permuted cell.
BatchNorm γ node = ∂G/∂γ, at the reassocB index the render's node reads.
BatchNorm β node = ∂G/∂β, at the reassocB index.
A channel-broadcast bias's Jacobian is the channel indicator.
A channel-broadcast bias, read by the β node, = ∂G/∂b. For any per-example op whose
bias enters as per 0 y + broadcast b, the emitted bnBetaGradB on the op's output cotangent
is the loss derivative in the bias.
A depthwise conv's bias is a channel broadcast.
A symmetric strided depthwise conv's bias is a channel broadcast: decimation keeps channels.
An XLA-SAME strided conv's bias is a channel broadcast.
The stride-4 patchify conv's bias is a channel broadcast: both decimations keep channels.