Documentation

LeanMlir.Proofs.Architectures.WindowMax

windowMax — the max pool over a family of product windows #

One max pool, proved once. Output cell (ch, hi, wi) is the max of the input over the window {(r hi a, s wi b) : a, b ∈ Fin k}, where r and s are the row and column index maps of the window. The pools the nets use are instances:

Smoothness is stated over positions #

WindowSmooth asks that a cell dominating its window be strictly above every cell at another input position, not at another offset. The two differ when r or s is not injective: the 3×3/s2 pool's clamped first window names one input cell at two offsets, and those two values are the same number, so a smoothness condition over offsets could never hold there. For the 2×2 pool offsets and positions coincide (windowSmooth_of_maxPool2Smooth).

Why overlapping windows cost nothing extra #

At a smooth point the pool is locally the reindexing y ↦ y ∘ σ (windowMax_flat_hasFDerivAt), σ sending each output to its argmax's input position. The argmax is the first maximal offset in row-major order (windowArgmax), the cell the emitted select_and_scatter (GE select) picks, so at a tie the gather and the printed scatter still name one cell. Overlapping windows only make σ non-injective, and reindexCLM's adjoint already sums over preimages, so the VJP (windowMaxHasVJPAt3) accumulates over every output whose window selects an input: one term for tiling windows, up to four for the 3×3/s2 pool.

noncomputable def Proofs.windowMax {c h w H W k : ℕ} [NeZero k] (r : Fin h → Fin k → Fin H) (s : Fin w → Fin k → Fin W) (x : Tensor3 c H W) :
Tensor3 c h w

The window max pool, [c, H, W] → [c, h, w]: the max of x ch over the window {(r hi a, s wi b)}.

Equations
Instances For
    theorem Proofs.le_windowMax {c h w H W k : ℕ} [NeZero k] (r : Fin h → Fin k → Fin H) (s : Fin w → Fin k → Fin W) (x : Tensor3 c H W) (ch : Fin c) (hi : Fin h) (wi : Fin w) (ab : Fin k × Fin k) :
    x ch (r hi ab.1) (s wi ab.2) ≤ windowMax r s x ch hi wi

    Every window cell is ≤ the pooled value.

    theorem Proofs.windowMax_attained {c h w H W k : ℕ} [NeZero k] (r : Fin h → Fin k → Fin H) (s : Fin w → Fin k → Fin W) (x : Tensor3 c H W) (ch : Fin c) (hi : Fin h) (wi : Fin w) :
    ∃ (ab : Fin k × Fin k), windowMax r s x ch hi wi = x ch (r hi ab.1) (s wi ab.2)

    The pooled value is attained by some window cell.

    theorem Proofs.windowMax_abs_le {c h w H W k : ℕ} [NeZero k] (r : Fin h → Fin k → Fin H) (s : Fin w → Fin k → Fin W) {x : Tensor3 c H W} {A : ℝ} (hx : ∀ (ci : Fin c) (hi : Fin H) (wi : Fin W), |x ci hi wi| ≤ A) (ci : Fin c) (hi : Fin h) (wi : Fin w) :
    |windowMax r s x ci hi wi| ≤ A

    The pool never grows magnitudes: it selects an existing window cell.

    theorem Proofs.windowMax_close {c h w H W k : ℕ} [NeZero k] (r : Fin h → Fin k → Fin H) (s : Fin w → Fin k → Fin W) (xt xa : Tensor3 c H W) {e : ℝ} (hx : ∀ (ci : Fin c) (hi : Fin H) (wi : Fin W), |xt ci hi wi - xa ci hi wi| ≤ e) (ci : Fin c) (hi : Fin h) (wi : Fin w) :
    |windowMax r s xt ci hi wi - windowMax r s xa ci hi wi| ≤ e

    The pool is 1-Lipschitz in the sup norm: each side is ≤ the other plus e, from the attained cell on one side and le_windowMax on the other.

    theorem Proofs.windowMax_shift {c h w H W k : ℕ} [NeZero k] (r : Fin h → Fin k → Fin H) (s : Fin w → Fin k → Fin W) (x y : Tensor3 c H W) (δ : ℝ) (ci : Fin c) (hxy : ∀ (p : Fin H) (q : Fin W), x ci p q = y ci p q + δ) (hi : Fin h) (wi : Fin w) :
    windowMax r s x ci hi wi = windowMax r s y ci hi wi + δ

    The pool shifts with a uniform offset. If one slab's channel is another's plus the constant δ, so are their pooled values (Finset.apply_sup'_eq_sup'_comp at (· + δ)). It holds at every point, with no argmax argument.

    theorem Proofs.windowMax_nonneg {c h w H W k : ℕ} [NeZero k] (r : Fin h → Fin k → Fin H) (s : Fin w → Fin k → Fin W) (x : Tensor3 c H W) (hx : ∀ (ci : Fin c) (p : Fin H) (q : Fin W), 0 ≤ x ci p q) (ci : Fin c) (hi : Fin h) (wi : Fin w) :
    0 ≤ windowMax r s x ci hi wi

    The pool keeps a nonnegative slab nonnegative.

    def Proofs.WindowSmooth {c h w H W k : ℕ} (r : Fin h → Fin k → Fin H) (s : Fin w → Fin k → Fin W) (x : Tensor3 c H W) :

    Smoothness: every window attains its max at exactly one input POSITION. A cell that dominates its window is strictly above every cell at another position; other cells may tie with each other. A window whose max sits at two positions does not qualify, and there the pool has no derivative.

    Equations
    • One or more equations did not get rendered due to their size.
    Instances For
      theorem Proofs.windowSmooth_of_pairwise {c h w H W k : ℕ} (r : Fin h → Fin k → Fin H) (s : Fin w → Fin k → Fin W) (x : Tensor3 c H W) (hd : ∀ (ci : Fin c) (hi_out : Fin h) (wi_out : Fin w) (ab ab' : Fin k × Fin k), (r hi_out ab.1, s wi_out ab.2) ≠ (r hi_out ab'.1, s wi_out ab'.2) → x ci (r hi_out ab.1) (s wi_out ab.2) ≠ x ci (r hi_out ab'.1) (s wi_out ab'.2)) :

      Windows whose cells at distinct positions have distinct values are smooth.

      theorem Proofs.windowSmooth_of_injective {c h w H W k : ℕ} (r : Fin h → Fin k → Fin H) (s : Fin w → Fin k → Fin W) (x : Tensor3 c H W) (hinj : ∀ (ci : Fin c) (p p' : Fin H) (q q' : Fin W), x ci p q = x ci p' q' → p = p' ∧ q = q') :

      Positional injectivity ⇒ smoothness. If on each channel (p, q) ↦ x ci p q is injective, no two positions tie. Because smoothness is quantified over positions, the injectivity lands directly on the hypothesis, with no decoding from positions back to offsets.

      def Proofs.WindowSmoothOrDead {c h w H W k : ℕ} (r : Fin h → Fin k → Fin H) (s : Fin w → Fin k → Fin W) (x : Tensor3 c H W) :

      Smooth or dead: every window either has its maximum at one input position, or has every cell ≤ 0. The second case is what a pool AFTER a ReLU needs: a window of dead ReLUs ties at 0, so the pool alone has no derivative there, but pool ∘ relu is locally constant at a pre-activation whose cells are all strictly negative.

      Equations
      • One or more equations did not get rendered due to their size.
      Instances For
        theorem Proofs.windowSmoothOrDead_of_smooth {c h w H W k : ℕ} (r : Fin h → Fin k → Fin H) (s : Fin w → Fin k → Fin W) {x : Tensor3 c H W} (hx : WindowSmooth r s x) :

        A smooth pool input is smooth-or-dead.

        def Proofs.WindowSmoothUpTo {c h w H W k : ℕ} (r : Fin h → Fin k → Fin H) (s : Fin w → Fin k → Fin W) (T : Fin H × Fin W → Fin H × Fin W → Prop) (x : Tensor3 c H W) :

        Smooth, dead, or tied only between twins: every window is entirely ≤ 0, or its maximum is strictly above every cell at another position except positions T relates to the maximum's. With T empty this is WindowSmoothOrDead. The twins a parameter gradient can afford are cells that are the SAME function of the moving parameter: a tie between them persists along the parameter, and the pool may pick either.

        Equations
        • One or more equations did not get rendered due to their size.
        Instances For
          theorem Proofs.windowSmoothUpTo_of_smoothOrDead {c h w H W k : ℕ} (r : Fin h → Fin k → Fin H) (s : Fin w → Fin k → Fin W) (T : Fin H × Fin W → Fin H × Fin W → Prop) {x : Tensor3 c H W} (hx : WindowSmoothOrDead r s x) :

          A smooth-or-dead pool input is smooth up to any twin relation.

          def Proofs.WindowMarginUpTo {c h w H W k : ℕ} (r : Fin h → Fin k → Fin H) (s : Fin w → Fin k → Fin W) (δ : ℝ) (T : Fin H × Fin W → Fin H × Fin W → Prop) (x : Tensor3 c H W) :

          Margin up to twins, the quantitative WindowSmoothUpTo: every window is entirely ≤ 0, or a cell dominating it is more than 2δ above every cell at another position, except positions T relates to its own. A perturbation of at most δ per entry then keeps every such cell strictly below, so the dominating cell keeps dominating; the twins it ties with must be the same function of whatever moves the input, and then they stay tied. Stated on the pool's input before any ReLU: the descent rungs apply it to the pre-activation, where a dead window is one whose cells are all ≤ 0.

          Equations
          • One or more equations did not get rendered due to their size.
          Instances For
            theorem Proofs.WindowMarginUpTo.mono {c h w H W k : ℕ} (r : Fin h → Fin k → Fin H) (s : Fin w → Fin k → Fin W) {δ δ' : ℝ} (hδ : δ' ≤ δ) {T : Fin H × Fin W → Fin H × Fin W → Prop} {x : Tensor3 c H W} (hx : WindowMarginUpTo r s δ T x) :
            WindowMarginUpTo r s δ' T x

            A margin up to twins holds at every smaller margin.

            theorem Proofs.windowMarginUpTo_of_cert {c h w H W k : ℕ} (r : Fin h → Fin k → Fin H) (s : Fin w → Fin k → Fin W) {δ : ℝ} (hδ : 0 ≤ δ) (T : Fin H × Fin W → Fin H × Fin W → Prop) (hsymm : ∀ (p q : Fin H × Fin W), T p q → T q p) (htrans : ∀ (p q u : Fin H × Fin W), T p q → T q u → T p u) {x : Tensor3 c H W} (hx : ∀ (ci : Fin c) (hi_out : Fin h) (wi_out : Fin w), (∀ (cd : Fin k × Fin k), x ci (r hi_out cd.1) (s wi_out cd.2) ≤ 0) ∨ ∃ (m : Fin k × Fin k), ∀ (cd : Fin k × Fin k), (r hi_out m.1, s wi_out m.2) = (r hi_out cd.1, s wi_out cd.2) ∨ T (r hi_out m.1, s wi_out m.2) (r hi_out cd.1, s wi_out cd.2) ∨ x ci (r hi_out cd.1) (s wi_out cd.2) + 2 * δ < x ci (r hi_out m.1) (s wi_out m.2)) :
            WindowMarginUpTo r s δ T x

            A margin up to twins from one designated cell per window. If T is an equivalence and every live window has a cell m with every other cell at m's position, a twin of it, or more than 2δ below it, the margin holds: a cell dominating the window is m or a twin of m (it cannot sit 2δ below), and twins of m inherit m's gaps. The form a concrete instance discharges, one certificate per window.

            theorem Proofs.windowSmoothUpTo_of_margin {c h w H W k : ℕ} (r : Fin h → Fin k → Fin H) (s : Fin w → Fin k → Fin W) {δ : ℝ} (hδ : 0 ≤ δ) (T : Fin H × Fin W → Fin H × Fin W → Prop) {x : Tensor3 c H W} (hx : WindowMarginUpTo r s δ T x) :

            At a nonnegative margin the margined input is smooth up to the same twins.

            noncomputable def Proofs.windowMaxOffsets {c h w H W k : ℕ} (r : Fin h → Fin k → Fin H) (s : Fin w → Fin k → Fin W) (x : Tensor3 c H W) (co : Fin c) (ho : Fin h) (wo : Fin w) :
            Finset (Fin (k * k))

            The offsets attaining the max of the window at output (co, ho, wo), as row-major flat indices a·k + b (finProdFinEquiv).

            Equations
            • One or more equations did not get rendered due to their size.
            Instances For
              theorem Proofs.windowMaxOffsets_nonempty {c h w H W k : ℕ} [NeZero k] (r : Fin h → Fin k → Fin H) (s : Fin w → Fin k → Fin W) (x : Tensor3 c H W) (co : Fin c) (ho : Fin h) (wo : Fin w) :
              (windowMaxOffsets r s x co ho wo).Nonempty
              noncomputable def Proofs.windowArgmax {c h w H W k : ℕ} [NeZero k] (r : Fin h → Fin k → Fin H) (s : Fin w → Fin k → Fin W) (x : Tensor3 c H W) (co : Fin c) (ho : Fin h) (wo : Fin w) :
              Fin k × Fin k

              The first argmax of the window at output (co, ho, wo), as an offset: the least maximal offset in row-major order (a major, windowArgmax_first). That is the cell the emitted select_and_scatter routes to: its GE select keeps the current pick while it is ≥ the next cell, so it ends on the first maximum in window iteration order. Unique under WindowSmooth up to position, and only windowArgmax_max is needed off a tie.

              Equations
              Instances For
                theorem Proofs.windowArgmax_max {c h w H W k : ℕ} [NeZero k] (r : Fin h → Fin k → Fin H) (s : Fin w → Fin k → Fin W) (x : Tensor3 c H W) (co : Fin c) (ho : Fin h) (wo : Fin w) (ab : Fin k × Fin k) :
                x co (r ho ab.1) (s wo ab.2) ≤ x co (r ho (windowArgmax r s x co ho wo).1) (s wo (windowArgmax r s x co ho wo).2)
                theorem Proofs.windowArgmax_first {c h w H W k : ℕ} [NeZero k] (r : Fin h → Fin k → Fin H) (s : Fin w → Fin k → Fin W) (x : Tensor3 c H W) (co : Fin c) (ho : Fin h) (wo : Fin w) (cd : Fin k × Fin k) (hcd : finProdFinEquiv cd < finProdFinEquiv (windowArgmax r s x co ho wo)) :
                x co (r ho cd.1) (s wo cd.2) < x co (r ho (windowArgmax r s x co ho wo).1) (s wo (windowArgmax r s x co ho wo).2)

                No earlier offset attains the max. Every offset before windowArgmax in row-major order is strictly below the window max: the first-maximum half of the select_and_scatter reading.

                theorem Proofs.windowMax_eq_at_max {c h w H W k : ℕ} [NeZero k] (r : Fin h → Fin k → Fin H) (s : Fin w → Fin k → Fin W) (x : Tensor3 c H W) (co : Fin c) (ho : Fin h) (wo : Fin w) (ab : Fin k × Fin k) (h_max : ∀ (cd : Fin k × Fin k), x co (r ho cd.1) (s wo cd.2) ≤ x co (r ho ab.1) (s wo ab.2)) :
                windowMax r s x co ho wo = x co (r ho ab.1) (s wo ab.2)

                If offset ab dominates every window cell, the pooled value is the value there.

                theorem Proofs.windowMax_eq_argmax_value {c h w H W k : ℕ} [NeZero k] (r : Fin h → Fin k → Fin H) (s : Fin w → Fin k → Fin W) (x : Tensor3 c H W) (co : Fin c) (ho : Fin h) (wo : Fin w) :
                windowMax r s x co ho wo = x co (r ho (windowArgmax r s x co ho wo).1) (s wo (windowArgmax r s x co ho wo).2)
                noncomputable def Proofs.windowGather {c h w H W k : ℕ} (r : Fin h → Fin k → Fin H) (s : Fin w → Fin k → Fin W) (σ : Fin c → Fin h → Fin w → Fin k × Fin k) (x : Tensor3 c H W) :
                Tensor3 c h w

                The window gather at a fixed selection σ: output (ch, hi, wi) reads the window cell at offset σ ch hi wi. Linear in x, with no argmax to decide. Wherever σ names a cell dominating every window, it IS the pool (windowMax_eq_windowGather); that is how a pool with tied windows is handled, the ties being routed to one fixed cell.

                Equations
                Instances For
                  theorem Proofs.windowMax_eq_windowGather {c h w H W k : ℕ} [NeZero k] (r : Fin h → Fin k → Fin H) (s : Fin w → Fin k → Fin W) (σ : Fin c → Fin h → Fin w → Fin k × Fin k) (x : Tensor3 c H W) (hdom : ∀ (ch : Fin c) (hi : Fin h) (wi : Fin w) (cd : Fin k × Fin k), x ch (r hi cd.1) (s wi cd.2) ≤ x ch (r hi (σ ch hi wi).1) (s wi (σ ch hi wi).2)) :
                  windowMax r s x = windowGather r s σ x

                  The pool is the gather at a dominating selection.

                  noncomputable def Proofs.windowLocalReindex {c h w H W k : ℕ} [NeZero k] (r : Fin h → Fin k → Fin H) (s : Fin w → Fin k → Fin W) (x : Tensor3 c H W) (k_out : Fin (c * h * w)) :
                  Fin (c * H * W)

                  For each output flat index, the flat index of its argmax's input position: the carrier of the local linearisation. Not injective when windows overlap; reindexCLM's adjoint sums over preimages, which is where the backward accumulates.

                  Equations
                  • One or more equations did not get rendered due to their size.
                  Instances For
                    theorem Proofs.windowMax_flat_hasFDerivAt {c h w H W k : ℕ} [NeZero k] (r : Fin h → Fin k → Fin H) (s : Fin w → Fin k → Fin W) (x : Tensor3 c H W) (h_smooth : WindowSmooth r s x) :

                    Smooth-point local linearisation. Near flatten x the flattened pool agrees with the reindex y ↦ y ∘ σ: every window keeps its argmax, since finitely many strict inequalities persist on a neighbourhood (Filter.eventually_all). An offset naming the argmax's own position is equal to it, not below, so the domination argument branches on positions.

                    theorem Proofs.pdiv3_windowMax_smooth {c h w H W k : ℕ} [NeZero k] (r : Fin h → Fin k → Fin H) (s : Fin w → Fin k → Fin W) (x : Tensor3 c H W) (h_smooth : WindowSmooth r s x) (ci : Fin c) (hi_in : Fin H) (wi_in : Fin W) (co : Fin c) (ho : Fin h) (wo : Fin w) :
                    pdiv3 (windowMax r s) x ci hi_in wi_in co ho wo = if windowLocalReindex r s x (finProdFinEquiv (finProdFinEquiv (co, ho), wo)) = finProdFinEquiv (finProdFinEquiv (ci, hi_in), wi_in) then 1 else 0

                    Smooth-point Jacobian. pdiv3 is the 0/1 indicator that the local reindex sends output (co, ho, wo) to input (ci, hi_in, wi_in). Left as the reindex equation: with overlapping windows an input has no single owning window to decode it into.

                    noncomputable def Proofs.windowMaxHasVJPAt3 {c h w H W k : ℕ} [NeZero k] (r : Fin h → Fin k → Fin H) (s : Fin w → Fin k → Fin W) (x : Tensor3 c H W) (h_smooth : WindowSmooth r s x) :

                    The VJP witness. The backward accumulates dy over every output whose window selects this input.

                    Equations
                    • One or more equations did not get rendered due to their size.
                    Instances For
                      noncomputable def Proofs.windowMaxFlat {c h w H W k : ℕ} [NeZero k] (r : Fin h → Fin k → Fin H) (s : Fin w → Fin k → Fin W) :
                      Vec (c * H * W) → Vec (c * h * w)

                      The flattened window max, the Vec-level form a codegen op denotes.

                      Equations
                      Instances For
                        theorem Proofs.windowMaxFlat_differentiableAt {c h w H W k : ℕ} [NeZero k] (r : Fin h → Fin k → Fin H) (s : Fin w → Fin k → Fin W) (x : Tensor3 c H W) (h_smooth : WindowSmooth r s x) :
                        theorem Proofs.windowMaxFlat_abs_le {c h w H W k : ℕ} [NeZero k] (r : Fin h → Fin k → Fin H) (s : Fin w → Fin k → Fin W) {v : Vec (c * H * W)} {A : ℝ} (hv : ∀ (i : Fin (c * H * W)), |v i| ≤ A) (i : Fin (c * h * w)) :

                        Flattened magnitude bound.

                        theorem Proofs.windowMaxFlat_close {c h w H W k : ℕ} [NeZero k] (r : Fin h → Fin k → Fin H) (s : Fin w → Fin k → Fin W) (vt va : Vec (c * H * W)) {e : ℝ} (hv : ∀ (i : Fin (c * H * W)), |vt i - va i| ≤ e) (i : Fin (c * h * w)) :
                        |windowMaxFlat r s vt i - windowMaxFlat r s va i| ≤ e

                        Flattened closeness.

                        theorem Proofs.windowMaxFlat_continuous {c h w H W k : ℕ} [NeZero k] (r : Fin h → Fin k → Fin H) (s : Fin w → Fin k → Fin W) :

                        windowMaxFlat is continuous (a sup' of coordinates).