A parametric lens is a morphism of , the Para Construction applied to the category of lenses over a cartesian category :
- objects: pairs ;
- morphisms : a parameter pair and a lens , i.e. two ordinary maps
Feed it a parameter, an input and a desired change in the output; get back a change in the parameter and a change in the input. Three wires in each direction: data left to right, corrections right to left, parameters top-down and parameter updates bottom-up. Composition joins the wires and leaves both parameter wires dangling, so the composite has parameter pair .
Sources: Cruttwell, Gavranović, Ghani, Wilson & Zanasi arXiv:2103.01931 (notes) Definition 2.5 and footnote 5 (these are the learners of Fong et al.); Fong, Spivak & Tuyéras, Backprop as Functor arXiv:1711.10455 (notes); Capucci et al. arXiv:2105.06332 (notes) (parametrised optics); Gavranović arXiv:2403.13001 (notes).
Why and not just
The update to a parameter need not live in the same space as the parameter. In they coincide () and this is the case everyone has in mind. But for Boolean circuits over the “gradient” is an XOR mask (Wilson & Zanasi, arXiv:2101.10488 (notes)); for a parameter constrained to a manifold — a rotation, a covariance matrix, a point of a Grassmannian — the update lives in a tangent (or cotangent) space; and turning a cotangent into a step needs a metric, which is where natural gradients enter.
Where parametric lenses come from, and what is still missing
- From autodiff. A reverse derivative turns a -map into a parametric lens: . A neural network layer with its
pullbackis exactly such an image. - Still missing for learning. A parametric lens takes a change in and returns a change in . When training one has a target value and wants a new parameter. The gaps are closed by a loss map, a learning rate and an optimiser — all of them lenses too; see Gradient-Based Learning with Parametric Lenses.
- Other backward passes. Replacing lenses by Bayesian lenses gives the parameterized statistical games of AutoBayes (the backward pass is a posterior); replacing them by optics with a selection functor gives open games (the backward pass is a best response).
Lenticulum.jl
Lux.jl layers are parametric lenses in ; Lenticulum.jl factors are parameterized statistical games, which only become a lens once a polarity is chosen. See Lux as a Parametric Lens.
Docs: plain Julia — Catlab has no dedicated API for this; related: Catlab v0.16 docs · GATlab standard library
# A parametric lens for a dense layer y = tanh.(W x): get(p, x), put(p, x, ȳ) = (p̄, x̄).
struct PLens{G,P}; get::G; put::P; end
dense_tanh = PLens(
(W, x) -> tanh.(W * x),
(W, x, ȳ) -> (δ = ȳ .* (1 .- tanh.(W * x) .^ 2); (δ * x', W' * δ)))
# composition: parameters pair up, the backward pass is the reverse chain rule
compose(g::PLens, f::PLens) = PLens(
((q, p), x) -> g.get(q, f.get(p, x)),
((q, p), x, z̄) -> begin
y = f.get(p, x)
q̄, ȳ = g.put(q, y, z̄)
p̄, x̄ = f.put(p, x, ȳ)
((q̄, p̄), x̄)
end)
net = compose(dense_tanh, dense_tanh)
W1, W2, x = [0.5 -0.2; 0.1 0.3], [1.0 0.4], [1.0, 2.0]
(ḡ2, ḡ1), x̄ = net.put((W2, W1), x, [1.0])
# check the first-layer gradient against finite differences
L(W1) = sum(tanh.(W2 * tanh.(W1 * x)))
e = [1.0 0; 0 0]; h = 1e-6
isapprox(ḡ1[1, 1], (L(W1 + h * e) - L(W1 - h * e)) / 2h; atol = 1e-6) # trueimport Mathlib
-- A parametric lens between (A, A') and (B, B') with parameter pair (P, P').
structure PLens (P P' A A' B B' : Type) where
get : P × A → B
put : P × A × B' → P' × A'
def PLens.comp {P P' Q Q' A A' B B' C C' : Type}
(f : PLens P P' A A' B B') (g : PLens Q Q' B B' C C') : PLens (Q × P) (Q' × P') A A' C C' where
get := fun ((q, p), a) => g.get (q, f.get (p, a))
put := fun ((q, p), a, c') =>
let (q', b') := g.put (q, f.get (p, a), c')
let (p', a') := f.put (p, a, b')
((q', p'), a')-- A parametric lens: get :: (p, a) -> b ; put :: (p, a, b') -> (p', a')
data PLens p p' a a' b b' = PLens { get :: (p, a) -> b, put :: (p, a, b') -> (p', a') }
(|>) :: PLens p p' a a' b b' -> PLens q q' b b' c c' -> PLens (q, p) (q', p') a a' c c'
PLens f f' |> PLens g g' = PLens
(\((q, p), a) -> g (q, f (p, a)))
(\((q, p), a, c') -> let (q', b') = g' (q, f (p, a), c')
(p', a') = f' (p, a, b')
in ((q', p'), a'))
-- a scalar linear layer y = w x
linear :: PLens Double Double Double Double Double Double
linear = PLens (\(w, x) -> w * x) (\(w, x, dy) -> (x * dy, w * dy))
-- put (linear |> linear) ((2, 3), 5, 1) == ((15, 10), 6)