Documentation

LeanMlir.Proofs.Architectures.GeluErf

The exact GELU: x · Φ(x) #

GELU as Hendrycks and Gimpel define it, gelu(x) = x · Φ(x) with Φ the standard normal CDF. This is the function PyTorch's nn.GELU and jax.nn.gelu(approximate=False) compute; Proofs.gelu in Activations is its tanh approximation.

Mathlib has no error function, so erf is defined here as (2/√π) ∫₀ᶻ exp(−t²) dt. The density and the CDF are closed forms over the interval integral; GeluErfGaussian proves them equal to Mathlib's gaussianPDFReal 0 1 and to the cdf of gaussianReal 0 1.

References #

noncomputable def Proofs.gaussPdf (x : ℝ) :

The standard normal density φ(x) = exp(−x²/2) / √(2π).

Equations
Instances For
    noncomputable def Proofs.gaussPhi (x : ℝ) :

    The standard normal CDF Φ(x) = ½ + ∫₀ˣ φ(t) dt.

    The density is even and integrates to one, so the mass below zero is ½ (integral_Iic_zero_gaussPdf); the interval integral carries the rest, and makes Φ' = φ the fundamental theorem of calculus (hasDerivAt_gaussPhi).

    Equations
    Instances For

      Φ' = φ — the fundamental theorem of calculus at the continuous density.

      gaussPhi is differentiable. Tagged for fun_prop so smoothness goals over the exact GELU dispatch.

      noncomputable def Proofs.geluErfScalar (x : ℝ) :

      Exact GELU forward — gelu(x) = x · Φ(x).

      Equations
      Instances For
        noncomputable def Proofs.geluErf (n : ℕ) (x : Vec n) :
        Vec n

        The elementwise exact GELU, applied componentwise to a vector.

        Equations
        Instances For
          noncomputable def Proofs.geluErfScalarDeriv (x : ℝ) :

          Scalar derivative of geluErfScalar — defined as Mathlib's deriv, as geluScalarDeriv is; its closed form is geluErfScalarDeriv_eq.

          Equations
          Instances For

            The product rule on x · Φ(x).

            Closed form of geluErfScalarDeriv — gelu'(x) = Φ(x) + x · φ(x).

            Differentiability of geluErf D as a function on Vec D.

            theorem Proofs.pdiv_geluErf (n : ℕ) (x : Vec n) (i j : Fin n) :
            pdiv (geluErf n) x i j = if i = j then geluErfScalarDeriv (x i) else 0

            Partial derivative of the exact GELU — diagonal, pdiv_elementwise at geluErfScalar.

            noncomputable def Proofs.geluErfHasVJP (n : ℕ) :

            Exact GELU VJP: elementwise multiply by the scalar derivative,

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

            Equations
            Instances For
              theorem Proofs.geluErfHasVJP_correct (n : ℕ) (x dy : Vec n) (i : Fin n) :
              (geluErfHasVJP n).backward x dy i = ∑ j : Fin n, pdiv (geluErf n) x i j * dy j

              Public correctness theorem for geluErfHasVJP: the exact GELU backward (diagonal scaling by geluErfScalarDeriv) equals the pdiv-contracted Jacobian.

              noncomputable def Proofs.erf (z : ℝ) :

              The error function erf(z) = (2/√π) ∫₀ᶻ exp(−t²) dt.

              Equations
              Instances For
                noncomputable def Proofs.erfc (z : ℝ) :

                The complementary error function erfc(z) = 1 − erf(z).

                Equations
                Instances For
                  theorem Proofs.erf_neg (z : ℝ) :
                  erf (-z) = -erf z

                  erf is odd: the integrand is even.

                  theorem Proofs.gaussPhi_eq_erf (x : ℝ) :
                  gaussPhi x = 1 / 2 * (1 + erf (x / √2))

                  Φ through erf — Φ(x) = ½ (1 + erf(x/√2)), by the substitution t = √2 · s.

                  theorem Proofs.gaussPhi_eq_erfc (x : ℝ) :
                  gaussPhi x = 0.5 * erfc (-x * √(1 / 2))

                  Φ through erfc — Φ(x) = ½ erfc(−x·√½). Equal to the erf spelling as real numbers (erf is odd); in floats this one keeps the negative tail, where 1 + erf cancels.

                  theorem Proofs.geluErfScalar_eq_erfc (x : ℝ) :
                  geluErfScalar x = 0.5 * x * erfc (-x * √(1 / 2))

                  The exact GELU forward as computed — gelu(x) = (0.5 · x) · erfc(−x · √½), the arithmetic of jax.nn.gelu(approximate=False).

                  theorem Proofs.geluErfScalarDeriv_eq_erfc (x : ℝ) :
                  geluErfScalarDeriv x = 0.5 * erfc (-x * √(1 / 2)) + 2 / √Real.pi * (0.5 * x) * Real.exp (-(-x * √(1 / 2)) ^ 2) * √(1 / 2)

                  The exact GELU derivative as computed — with z = −x · √½,

                  gelu'(x) = 0.5 · erfc(z) + (2/√π) · (0.5 · x) · exp(−z²) · √½,

                  the two terms jax.vjp of jax.nn.gelu(approximate=False) forms: Φ(x) through erfc, and x · φ(x) with the density written as the derivative of erfc at z.