The 8-conv CIFAR CNN with per-channel BN — every gradient node IS the loss's derivative #
cifar8Bn_train_step_tiedG states the 38 un-fused gradient nodes the packed cifar8w_bn_* arms
emit, each at the cotangent the chain threads to it. cifar8Bn_net_lossGrad states that each node,
at the chain cotangent, is the gradient of the loss in that parameter, for any loss L of the
logits with gradient g there; cifar8Bn_net_lossGrad_CE instantiates it at the softmax
cross-entropy the render emits.
The net is the BN-free 8-conv net (Cifar8TieG.cifar8_net_lossGrad) with a per-example,
per-channel BN between each conv and its ReLU, so each pool's pre-activation is a BN output. Two
cells equal at every conv weight stay equal through BN (one affine map per channel), so the twin
relations (Cifar8BnPoolTwin1 … Cifar8BnPoolTwin4, cells equal at every weight upstream of the
pool, BN γ/β included) and the selection routing carry over unchanged. Per node kind, BN adds
bnGamma_hasGradAt, bnBeta_hasGradAt and the input pull-back hasGradAt_bnPC.
Hypotheses. Odd kernels, every BN ε > 0 (Cifar8BnPos), every ReLU off its kink, every pool
window dead or tied only between twins, each selection naming a maximum of every window
(Cifar8BnLossSmoothAt).
Scope. One example (the emitted module batch-contracts; den is per-example; BN normalises
each channel over the example's own spatial cells).
Per-channel BN's VJP at a point: the renderable backward bnPerChannelTensor3GradInput.
Equations
- Proofs.Cifar8BnTieG.bnPCHasVJPAt oc h w ε hε γ β v = { backward := fun (dy : Proofs.Vec (oc * h * w)) => Proofs.bnPerChannelTensor3GradInput oc h w ε γ v dy, correct := ⋯ }
Instances For
Through per-channel BN: the backward is bnPerChannelTensor3GradInput, the BN-back the
render emits.
BN γ node = ∇_γ G.
BN β node = ∇_β G.
From one pool's pre-activation to the next pool's: ReLU, pool, conv, BN, ReLU, conv, BN.
Equations
- One or more equations did not get rendered due to their size.
Instances For
The first pool's pre-activation (BN₂'s output).
Equations
- One or more equations did not get rendered due to their size.
Instances For
Pool 2's pre-activation (BN₄'s output).
Equations
- One or more equations did not get rendered due to their size.
Instances For
Pool 3's pre-activation (BN₆'s output).
Equations
- One or more equations did not get rendered due to their size.
Instances For
Pool 4's pre-activation (BN₈'s output).
Equations
- One or more equations did not get rendered due to their size.
Instances For
Twins of pool 1: equal at every weight (conv and BN γ/β) upstream of it, in
every channel.
Equations
- One or more equations did not get rendered due to their size.
Instances For
Twins of pool 2: equal at every weight (conv and BN γ/β) upstream of it, in
every channel.
Equations
- One or more equations did not get rendered due to their size.
Instances For
Twins of pool 3: equal at every weight (conv and BN γ/β) upstream of it, in
every channel.
Equations
- One or more equations did not get rendered due to their size.
Instances For
Twins of pool 4: equal at every weight (conv and BN γ/β) upstream of it, in
every channel.
Equations
- One or more equations did not get rendered due to their size.
Instances For
The smooth-point bundle the loss gradient needs. Every ReLU off its kink (at the BN outputs and the dense head); every window of each pool dead or tied only between that pool's twins; each selection naming a maximum of every window.
- pool1 : SmallParamGrad.MaxPool2SmoothUpTo (Cifar8BnPoolTwin1 c1 kH kW ε₁ ε₂ x) (Tensor3.unflatten (cifar8BnPre1 W₁ b₁ ε₁ γ₁ β₁ W₂ b₂ ε₂ γ₂ β₂ x))
- z3 (k : Fin (c2 * (2 * (2 * (2 * h))) * (2 * (2 * (2 * w))))) : bnPerChannelTensor3 c2 (2 * (2 * (2 * h))) (2 * (2 * (2 * w))) ε₃ γ₃ β₃ (flatConv W₃ b₃ (maxPoolFlat c1 (2 * (2 * (2 * h))) (2 * (2 * (2 * w))) (relu (c1 * (2 * (2 * (2 * (2 * h)))) * (2 * (2 * (2 * (2 * w))))) (cifar8BnPre1 W₁ b₁ ε₁ γ₁ β₁ W₂ b₂ ε₂ γ₂ β₂ x)))) k ≠ 0
- pool2 : SmallParamGrad.MaxPool2SmoothUpTo (Cifar8BnPoolTwin2 c1 c2 kH kW ε₁ ε₂ ε₃ ε₄ x) (Tensor3.unflatten (cifar8BnPre2 W₁ b₁ ε₁ γ₁ β₁ W₂ b₂ ε₂ γ₂ β₂ W₃ b₃ ε₃ γ₃ β₃ W₄ b₄ ε₄ γ₄ β₄ x))
- z5 (k : Fin (c3 * (2 * (2 * h)) * (2 * (2 * w)))) : bnPerChannelTensor3 c3 (2 * (2 * h)) (2 * (2 * w)) ε₅ γ₅ β₅ (flatConv W₅ b₅ (maxPoolFlat c2 (2 * (2 * h)) (2 * (2 * w)) (relu (c2 * (2 * (2 * (2 * h))) * (2 * (2 * (2 * w)))) (cifar8BnPre2 W₁ b₁ ε₁ γ₁ β₁ W₂ b₂ ε₂ γ₂ β₂ W₃ b₃ ε₃ γ₃ β₃ W₄ b₄ ε₄ γ₄ β₄ x)))) k ≠ 0
- pool3 : SmallParamGrad.MaxPool2SmoothUpTo (Cifar8BnPoolTwin3 c1 c2 c3 kH kW ε₁ ε₂ ε₃ ε₄ ε₅ ε₆ x) (Tensor3.unflatten (cifar8BnPre3 W₁ b₁ ε₁ γ₁ β₁ W₂ b₂ ε₂ γ₂ β₂ W₃ b₃ ε₃ γ₃ β₃ W₄ b₄ ε₄ γ₄ β₄ W₅ b₅ ε₅ γ₅ β₅ W₆ b₆ ε₆ γ₆ β₆ x))
- sel3 : SmallParamGrad.PoolSelDom σ₃ (relu (c3 * (2 * (2 * h)) * (2 * (2 * w))) (cifar8BnPre3 W₁ b₁ ε₁ γ₁ β₁ W₂ b₂ ε₂ γ₂ β₂ W₃ b₃ ε₃ γ₃ β₃ W₄ b₄ ε₄ γ₄ β₄ W₅ b₅ ε₅ γ₅ β₅ W₆ b₆ ε₆ γ₆ β₆ x))
- pool4 : SmallParamGrad.MaxPool2SmoothUpTo (Cifar8BnPoolTwin4 c1 c2 c3 c4 kH kW ε₁ ε₂ ε₃ ε₄ ε₅ ε₆ ε₇ ε₈ x) (Tensor3.unflatten (cifar8BnPre4 W₁ b₁ ε₁ γ₁ β₁ W₂ b₂ ε₂ γ₂ β₂ W₃ b₃ ε₃ γ₃ β₃ W₄ b₄ ε₄ γ₄ β₄ W₅ b₅ ε₅ γ₅ β₅ W₆ b₆ ε₆ γ₆ β₆ W₇ b₇ ε₇ γ₇ β₇ W₈ b₈ ε₈ γ₈ β₈ x))
- sel4 : SmallParamGrad.PoolSelDom σ₄ (relu (c4 * (2 * h) * (2 * w)) (cifar8BnPre4 W₁ b₁ ε₁ γ₁ β₁ W₂ b₂ ε₂ γ₂ β₂ W₃ b₃ ε₃ γ₃ β₃ W₄ b₄ ε₄ γ₄ β₄ W₅ b₅ ε₅ γ₅ β₅ W₆ b₆ ε₆ γ₆ β₆ W₇ b₇ ε₇ γ₇ β₇ W₈ b₈ ε₈ γ₈ β₈ x))
- z9 (k : Fin d1) : dense W₉ b₉ (maxPoolFlat c4 h w (relu (c4 * (2 * h) * (2 * w)) (cifar8BnPre4 W₁ b₁ ε₁ γ₁ β₁ W₂ b₂ ε₂ γ₂ β₂ W₃ b₃ ε₃ γ₃ β₃ W₄ b₄ ε₄ γ₄ β₄ W₅ b₅ ε₅ γ₅ β₅ W₆ b₆ ε₆ γ₆ β₆ W₇ b₇ ε₇ γ₇ β₇ W₈ b₈ ε₈ γ₈ β₈ x))) k ≠ 0
Instances For
Every cifar8-bn gradient node is the gradient of L in that parameter: the 38 un-fused
nodes cifar8Bn_train_step_tiedG states, each at the cotangent the chain threads to its layer
(each pool routed at its selection), stated against L of cifarCnnBn8Forward with that one
parameter varied (F is L of the forward at the given weights, the εs fixed).
Equations
- One or more equations did not get rendered due to their size.
Instances For
Every cifar8-bn gradient node is the gradient of L in that parameter, whenever g is
L's gradient at the logits.
Hypotheses: odd kernels, every BN ε positive (Cifar8BnPos), and Cifar8BnLossSmoothAt —
every ReLU off its kink, every window of each pool dead or tied only between cells that are
the same function of the weights upstream of it, each selection naming a maximum of every
window.
The artifact's loss: every node is the gradient of the softmax cross-entropy at label,
g the emitted loss cotangent.