Documentation

LeanMlir.Proofs.Architectures.GeluForm

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.

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.

Which GELU: the tanh approximation or the exact x · Φ(x).

  • tanh : GeluForm

    0.5 · x · (1 + tanh(√(2/π) · (x + 0.044715 · x³))), geluScalar.

  • erf : GeluForm

    x · Φ(x), geluErfScalar.

Instances For
    @[instance_reducible]
    Equations
    noncomputable def Proofs.GeluForm.scalar :
    GeluForm → ℝ → ℝ

    The scalar activation of a form.

    Equations
    Instances For

      Either form is differentiable.

      noncomputable def Proofs.GeluForm.scalarDeriv (gf : GeluForm) (x : ℝ) :

      Scalar derivative of a form — Mathlib's deriv, as geluScalarDeriv and geluErfScalarDeriv are.

      Equations
      Instances For
        noncomputable def Proofs.GeluForm.map (gf : GeluForm) (n : ℕ) (x : Vec n) :
        Vec n

        The GELU of a form, applied componentwise to a vector.

        Equations
        Instances For

          Differentiability of gf.map D as a function on Vec D.

          theorem Proofs.GeluForm.pdiv_map (gf : GeluForm) (n : ℕ) (x : Vec n) (i j : Fin n) :
          pdiv (gf.map n) x i j = if i = j then gf.scalarDeriv (x i) else 0

          Partial derivative of either GELU — diagonal, pdiv_elementwise at gf.scalar.

          noncomputable def Proofs.GeluForm.hasVJP (gf : GeluForm) (n : ℕ) :
          HasVJP (gf.map n)

          VJP of either GELU: elementwise multiply by the scalar derivative,

          back(x, dy)_i = dy_i * gf.scalarDeriv(x_i).

          Equations
          Instances For
            theorem Proofs.GeluForm.hasVJP_correct (gf : GeluForm) (n : ℕ) (x dy : Vec n) (i : Fin n) :
            (gf.hasVJP n).backward x dy i = ∑ j : Fin n, pdiv (gf.map n) x i j * dy j

            Public correctness theorem for GeluForm.hasVJP: the backward (diagonal scaling by gf.scalarDeriv) equals the pdiv-contracted Jacobian.