Documentation

LeanMlir.Proofs.Architectures.Activations

Smooth activations: GELU, Swish, sigmoid #

The smooth elementwise activations the nets use, each with its scalar derivative, the closed form the emitted backward computes, and its VJP:

GELU is gelu(x) = x · Phi(x), where Phi is the CDF of the standard normal; in practice everyone uses the tanh approximation gelu(x) ~ 0.5 x (1 + tanh(sqrt(2/pi)(x + 0.044715 x^3))), because it is faster than the exact erf form, and that is the function here and in the renders.

Each Jacobian is diagonal (pdiv_elementwise), so each VJP is one line. ReLU and ReLU6, which have kinks, live in Foundation.MLP.

References #

noncomputable def Proofs.geluScalar (x : ℝ) :

GELU forward — Gaussian Error Linear Unit, tanh approximation.

gelu(x) = 0.5 · x · (1 + tanh(√(2/π) · (x + 0.044715 · x³)))

Matches the MLIR codegen (which emits the tanh approximation rather than the exact x · Φ(x) erf form).

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

    The elementwise GELU, applied componentwise to a vector.

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

      Scalar derivative of geluScalar — defined as Mathlib's deriv.

      geluScalar is the tanh approximation, so this is the derivative of that approximation; its closed form is geluScalarDeriv_eq. We define it via deriv rather than writing the closed form so the connection to geluScalar is automatic.

      Equations
      Instances For

        Real.tanh is differentiable everywhere — bridge via Real.tanh_eq_sinh_div_cosh and Real.cosh_pos. Tagged for fun_prop so downstream gelu-style smoothness goals dispatch. (Mathlib has neither this nor hasDerivAt_tanh below for Real.tanh.)

        Derivative of Real.tanh — tanh'(y) = 1 − tanh²(y), built from tanh = sinh/cosh via the quotient rule and cosh² − sinh² = 1, for the GELU closed-form derivative geluScalarDeriv_eq.

        theorem Proofs.geluScalarDeriv_eq (x : ℝ) :
        geluScalarDeriv x = 0.5 * (1 + Real.tanh (√(2 / Real.pi) * (x + 44715e-6 * x ^ 3))) + 0.5 * x * ((1 - Real.tanh (√(2 / Real.pi) * (x + 44715e-6 * x ^ 3)) ^ 2) * (√(2 / Real.pi) * (1 + 44715e-6 * (3 * x ^ 2))))

        Closed form of geluScalarDeriv — the analytic derivative of the tanh-approximation GELU. With u = √(2/π)·(x + 0.044715·x³) and t = tanh u,

        gelu'(x) = 0.5·(1 + t) + 0.5·x·(1 − t²)·√(2/π)·(1 + 3·0.044715·x²).

        This is exactly the closed form the verified geluBack StableHLO emitter renders — so the emitted backward text is certified equal to deriv geluScalar (swishScalarDeriv_eq does the same for swish). Proof: assemble HasDerivAt for the polynomial inner, tanh via hasDerivAt_tanh, and the outer product, then HasDerivAt.deriv.

        Differentiability of geluScalar as a scalar function.

        Differentiability of gelu D as a function on Vec D.

        theorem Proofs.pdiv_gelu (n : ℕ) (x : Vec n) (i j : Fin n) :
        pdiv (gelu n) x i j = if i = j then geluScalarDeriv (x i) else 0

        Partial derivative of GELU.

        gelu n has diagonal Jacobian: each output coord depends only on the corresponding input coord via geluScalar. So ∂(gelu n y)_j / ∂y_i = (geluScalar' (y i)) if i = j, else 0 — pdiv_elementwise at geluScalar.

        noncomputable def Proofs.geluHasVJP (n : ℕ) :

        GELU VJP: elementwise multiply by the scalar derivative.

        back(x, dy)_i = dy_i * geluScalarDeriv(x_i)

        Same template as ReLU (reluHasVJP), Swish, h-swish. If your activation has a diagonal Jacobian, this is the only proof you need — "collapse the diagonal sum."

        Equations
        Instances For
          theorem Proofs.geluHasVJP_correct (n : ℕ) (x dy : Vec n) (i : Fin n) :
          (geluHasVJP n).backward x dy i = ∑ j : Fin n, pdiv (gelu n) x i j * dy j

          Public correctness theorem for geluHasVJP: the GELU backward (diagonal scaling by geluScalarDeriv) equals the pdiv-contracted Jacobian.

          The activation taxonomy is closed #

          Every activation function in every architecture in this repo is elementwise -> diagonal Jacobian -> one-line VJP. Taking inventory:

          Activationpdiv_* formula (at j = i)
          ReLU1 if x_i > 0, else 0
          ReLU61 if 0 < x_i < 6, else 0
          Swishsigma(x_i) * (1 + x_i * (1 - sigma(x_i)))
          h-swishpiecewise: 0 / (2x_i + 3)/6 / 1
          h-sigmoidpiecewise: 0 / 1/6 / 0
          GELU (tanh approx.)geluScalarDeriv_eq
          tanh1 - tanh^2(x_i)
          sigmoidsigma(x_i) * (1 - sigma(x_i))

          They all have the same proof shape (pdiv_elementwise, then collapse the diagonal sum). This file instantiates it as geluHasVJP, swishHasVJP and sigmoidHasVJP; ReLU's and ReLU6's (reluHasVJP, relu6HasVJPAt) are in Foundation.MLP, and the other rows are not given separate HasVJP instances.

          Swish (a.k.a. SiLU) #

          swish(x) = x * σ(x), where σ(x) = 1 / (1 + exp(-x)) is the standard logistic sigmoid. Used as the default activation in EfficientNet's MBConv blocks. Same diagonal-Jacobian proof template as ReLU and GELU.

          noncomputable def Proofs.swishScalar (x : ℝ) :

          Swish forward — Sigmoid-Linear Unit (SiLU).

          swish(x) = x / (1 + exp(-x)) = x · σ(x). Smooth everywhere (denominator is bounded below by 1 > 0).

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

            The elementwise Swish, applied componentwise to a vector.

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

              Scalar derivative of swishScalar — defined via Mathlib's deriv. The closed form is σ(x)·(1 + x·(1 - σ(x))) (swishScalarDeriv_eq); we define it as deriv swishScalar so the link to swishScalar is automatic.

              Equations
              Instances For

                swishScalar x = x · σ(x) with Mathlib's logistic function Real.sigmoid.

                Closed form of swishScalarDeriv: σ(x)·(1 + x·(1 − σ(x))), from Real.hasDerivAt_sigmoid — the formula the swishBack StableHLO emitter renders.

                Differentiability of swishScalar. The denominator 1 + exp(-x) is always positive, so the quotient is smooth everywhere.

                Differentiability of swish D as a function on Vec D.

                theorem Proofs.pdiv_swish (n : ℕ) (x : Vec n) (i j : Fin n) :
                pdiv (swish n) x i j = if i = j then swishScalarDeriv (x i) else 0

                Partial derivative of Swish — diagonal Jacobian: each output coord depends only on the corresponding input coord via swishScalar (pdiv_elementwise).

                noncomputable def Proofs.swishHasVJP (n : ℕ) :

                Swish VJP: elementwise multiply by the scalar derivative. Same template as ReLU/GELU. The codegen emits the closed-form σ(x)·(1 + x·(1 - σ(x))) directly; swishScalarDeriv_eq equates it with swishScalarDeriv = deriv swishScalar.

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

                  The logistic function 1 / (1 + e^{−x}) on one real.

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

                    sigmoidScalar applied to each entry of a vector.

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

                      The derivative of sigmoidScalar at x, as Mathlib's deriv.

                      Equations
                      Instances For

                        The closed form σ' = σ·(1 − σ). sigmoidHasVJP's backward is stated with deriv; the emitted sigmoidBack text computes σ(x)·(1 − σ(x)) — this is the equation between them, so it stays pinned in the axiom audit although no Lean proof consumes it.

                        theorem Proofs.pdiv_sigmoid (n : ℕ) (x : Vec n) (i j : Fin n) :
                        pdiv (sigmoid n) x i j = if i = j then sigmoidScalarDeriv (x i) else 0
                        noncomputable def Proofs.sigmoidHasVJP (n : ℕ) :

                        The VJP of elementwise sigmoid: dy ⊙ σ'(x).

                        Equations
                        Instances For