theorem definition example

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
end
import 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))