Fong, Spivak and Tuyéras define a learner as a tuple of a parameter set and functions
taken up to relabelling of parameters. Learners compose — the request of the second learner is the training target passed back to the first — and form a symmetric monoidal category (Proposition II.4). The main theorem:
Theorem III.2. Fix a step size and a differentiable error function with invertible for each . Then gradient descent and backpropagation define a faithful, injective-on-objects, strong symmetric monoidal functor sending a parametrised function to the learner with and , where .
Backpropagation is functorial: training a composite network by gradient descent is the same as composing the learners obtained from its layers. The request map is the new ingredient — it is what a layer asks its predecessor to output instead.
Sources: Fong, Spivak & Tuyéras, Backprop as Functor: A compositional perspective on supervised learning arXiv:1711.10455 (notes) Definitions II.1, III.1, Proposition II.4, Theorem III.2; Fong & Johnson, Lenses and Learners arXiv:1903.03671 (notes); Cruttwell et al. arXiv:2103.01931 (notes) footnote 5, §6.
Learners are parametric lenses
Pair the update and request into one map and the learner is a lens with a parameter — a Parametric Lens, with the backward direction carrying targets rather than gradients. Fong & Johnson make this precise: learners embed in a category of (asymmetric) lenses, and the lens laws correspond to well-behaved learning. Cruttwell et al. (footnote 5) note that their parametric lenses are these learners, generalised from to any Reverse Derivative Category and from a fixed loss and step size to arbitrary loss maps, learning rates and optimisers (Gradient-Based Learning with Parametric Lenses).
What the theorem buys
- Modularity: a neural network’s learner is determined by its layers’ learners; the global learning rule need never be written down.
- Wiring diagrams: because is monoidal, the string diagram of a network — layers in series and in parallel, with copying and summing of wires — can be read as a diagram in .
- A caveat recorded in the paper: the invertibility condition on holds for quadratic error but not for every loss; later formulations (Cruttwell et al.) avoid it by separating the loss map from the learning rate.
Docs: plain Julia — Catlab has no dedicated API for this; related: Catlab v0.16 docs · GATlab standard library
# A learner (I, U, r) and its composite; check that training a composite equals
# composing the learners of the parts (Theorem III.2 in miniature, quadratic error).
struct Learner{I,U,R}; I::I; U::U; r::R; end
const ε = 0.05
function gd_learner(I, ∇p, ∇a) # ∇p, ∇a of E(p,a,b) = ½(I(p,a) - b)²
Learner(I, (p, a, b) -> p - ε * ∇p(p, a, b), (p, a, b) -> a - ∇a(p, a, b))
end
lin = gd_learner((p, a) -> p * a, (p, a, b) -> (p * a - b) * a, (p, a, b) -> (p * a - b) * p)
function compose(L2::Learner, L1::Learner)
Learner(((q, p), a) -> L2.I(q, L1.I(p, a)),
((q, p), a, c) -> (b = L1.I(p, a); (L2.U(q, b, c), L1.U(p, a, L2.r(q, b, c)))),
((q, p), a, c) -> L1.r(p, a, L2.r(q, L1.I(p, a), c)))
end
net = compose(lin, lin)
let θ = (1.0, 0.5)
for _ in 1:2000, a in (-1.0, 0.5, 2.0); θ = net.U(θ, a, 6a); end # learn a ↦ 6a
round(net.I(θ, 1.0); digits = 2) # ≈ 6.0
endimport Mathlib
-- Fong–Spivak–Tuyéras learners (Definition II.1)
structure Learner (A B : Type) where
P : Type
implement : P × A → B
update : P × A × B → P
request : P × A × B → A
def Learner.comp {A B C : Type} (L₁ : Learner A B) (L₂ : Learner B C) : Learner A C where
P := L₂.P × L₁.P
implement := fun ((q, p), a) => L₂.implement (q, L₁.implement (p, a))
update := fun ((q, p), a, c) =>
let b := L₁.implement (p, a)
(L₂.update (q, b, c), L₁.update (p, a, L₂.request (q, b, c)))
request := fun ((q, p), a, c) => L₁.request (p, a, L₂.request (q, L₁.implement (p, a), c)){-# LANGUAGE ExistentialQuantification #-}
-- A learner a -> b with a hidden parameter type p (Definition II.1)
data Learner a b = forall p. Learner p (p -> a -> b) (p -> a -> b -> p) (p -> a -> b -> a)
compose :: Learner b c -> Learner a b -> Learner a c
compose (Learner q i2 u2 r2) (Learner p i1 u1 r1) =
Learner (q, p)
(\(q', p') a -> i2 q' (i1 p' a))
(\(q', p') a c -> let b = i1 p' a in (u2 q' b c, u1 p' a (r2 q' b c)))
(\(q', p') a c -> r1 p' a (r2 q' (i1 p' a) c))