Either GELU #
The GELU comes in two forms: the tanh approximation (geluScalar, the default of
jax.nn.gelu) and the exact x · Φ(x) (geluErfScalar, PyTorch's nn.GELU). GeluForm names
the choice. A net that uses the GELU, and every theorem about it, takes a GeluForm and holds
for both, so the renders that trained under the tanh form and the ones that train under the
exact form are instances of one statement.
GeluForm.scalar,GeluForm.scalarDeriv: the scalar activation and its derivative (aderiv; the closed forms aregeluScalarDeriv_eqandgeluErfScalarDeriv_eq).GeluForm.map,GeluForm.pdiv_map,GeluForm.hasVJP: the activation on a vector, its diagonal Jacobian, its VJP.GeluForm.map_tanh,GeluForm.map_erf,GeluForm.hasVJP_tanh,GeluForm.hasVJP_erf: at each form these aregelu/geluHasVJPandgeluErf/geluErfHasVJP, byrfl.
Nothing here cases on the form except scalar and its differentiability, so a definition built
on GeluForm.map unfolds the same way at either form.
Equations
- Proofs.instReprGeluForm.repr Proofs.GeluForm.tanh prec✝ = Repr.addAppParen (Std.Format.nest (if prec✝ ≥ 1024 then 1 else 2) (Std.Format.text "Proofs.GeluForm.tanh")).group prec✝
- Proofs.instReprGeluForm.repr Proofs.GeluForm.erf prec✝ = Repr.addAppParen (Std.Format.nest (if prec✝ ≥ 1024 then 1 else 2) (Std.Format.text "Proofs.GeluForm.erf")).group prec✝
Instances For
Equations
- Proofs.instReprGeluForm = { reprPrec := Proofs.instReprGeluForm.repr }
Equations
The scalar activation of a form.
Equations
Instances For
Either form is differentiable.
Scalar derivative of a form — Mathlib's deriv, as geluScalarDeriv and
geluErfScalarDeriv are.
Equations
- gf.scalarDeriv x = deriv gf.scalar x
Instances For
Differentiability of gf.map D as a function on Vec D.
VJP of either GELU: elementwise multiply by the scalar derivative,
back(x, dy)_i = dy_i * gf.scalarDeriv(x_i).
Equations
- gf.hasVJP n = { backward := fun (x dy : Proofs.Vec n) (i : Fin n) => dy i * gf.scalarDeriv (x i), correct := ⋯ }
Instances For
Public correctness theorem for GeluForm.hasVJP: the backward (diagonal scaling by
gf.scalarDeriv) equals the pdiv-contracted Jacobian.