Documentation

LeanMlir.Proofs.Foundation.PairTile

Pair-tile VJP #

Layer.pairTile (the distogram stem): two residue blocks Xi, Xj : Mat L D — the host's features, constants of the step — and two weights W, Wj : Mat D C; the pair map is y i j c = (Xi · W) i c + (Xj · Wj) j c. It is linear in each weight, so each weight's VJP is the adjoint of a linear map: dW = Xiᵀ · (Σ_j dy) and dWj = Xjᵀ · (Σ_i dy) — the cotangent summed over the axis the weight was broadcast along, then the dense weight rule. These are the two stablehlo.reduce + dot_general pairs the pairTile backward emits (%ptb_du / %d_W, %ptb_dv / %d_Wj). The layer takes no input gradient: its input is the feature block.

Both witnesses are pdiv_of_affine with the other block's term as the constant, the shape of pdiv_dense_W; the only work is the row-major (i, (j, c)) index of the flat pair map.

@[reducible]
def Proofs.PairTile.oidx {L C : ℕ} (o : Fin (L * (L * C))) :
Fin L × Fin L × Fin C

The (i, j, c) behind a flat index of a row-major [L, L, C] pair map.

Equations
Instances For
    @[simp]
    theorem Proofs.PairTile.oidx_mk {L C : ℕ} (a b : Fin L) (c : Fin C) :
    theorem Proofs.PairTile.sum_finProdFinEquiv_r {M : Type u_1} [AddCommMonoid M] {a b c : ℕ} (f : Fin (a * (b * c)) → M) :
    ∑ k : Fin (a * (b * c)), f k = ∑ i : Fin a, ∑ j : Fin b, ∑ l : Fin c, f (finProdFinEquiv (i, finProdFinEquiv (j, l)))

    A Fin (a·(b·c)) sum is the row-major triple sum, the pair map's (i, (j, c)) layout (sum_finProdFinEquiv₃ is the left-nested (a·b)·c).

    noncomputable def Proofs.PairTile.tileW {L D C : ℕ} (Xi Xj : Mat L D) (Wj : Mat D C) (v : Vec (D * C)) :
    Vec (L * (L * C))

    The pair map as a function of the flattened i-block weight (the j term is constant).

    Equations
    • One or more equations did not get rendered due to their size.
    Instances For
      noncomputable def Proofs.PairTile.tileWj {L D C : ℕ} (Xi Xj : Mat L D) (W : Mat D C) (v : Vec (D * C)) :
      Vec (L * (L * C))

      The pair map as a function of the flattened j-block weight (the i term is constant).

      Equations
      Instances For
        noncomputable def Proofs.PairTile.gradW {L D C : ℕ} (Xi : Mat L D) (dy : Vec (L * (L * C))) :
        Vec (D * C)

        dW (d, c) = Σ_i Xi i d · Σ_j dy i j c — the cotangent summed over j, contracted with Xi over i (the emitted reduce … dimensions = [2] then dot_general … [0, 1] x [0, 1], per sample).

        Equations
        Instances For
          noncomputable def Proofs.PairTile.gradWj {L D C : ℕ} (Xj : Mat L D) (dy : Vec (L * (L * C))) :
          Vec (D * C)

          dWj (d, c) = Σ_j Xj j d · Σ_i dy i j c — summed over i, contracted with Xj over j.

          Equations
          Instances For
            theorem Proofs.PairTile.pdiv_tileW {L D C : ℕ} (Xi Xj : Mat L D) (Wj : Mat D C) (v : Vec (D * C)) (k : Fin (D * C)) (o : Fin (L * (L * C))) :
            pdiv (tileW Xi Xj Wj) v k o = if (oidx o).2.2 = (finProdFinEquiv.symm k).2 then Xi (oidx o).1 (finProdFinEquiv.symm k).1 else 0

            ∂ y_{i j c} / ∂ W_{d c'} = Xi i d · δ(c, c').

            theorem Proofs.PairTile.pdiv_tileWj {L D C : ℕ} (Xi Xj : Mat L D) (W : Mat D C) (v : Vec (D * C)) (k : Fin (D * C)) (o : Fin (L * (L * C))) :
            pdiv (tileWj Xi Xj W) v k o = if (oidx o).2.2 = (finProdFinEquiv.symm k).2 then Xj (oidx o).2.1 (finProdFinEquiv.symm k).1 else 0

            ∂ y_{i j c} / ∂ Wj_{d c'} = Xj j d · δ(c, c').

            noncomputable def Proofs.PairTile.tileWHasVJP {L D C : ℕ} (Xi Xj : Mat L D) (Wj : Mat D C) :
            HasVJP (tileW Xi Xj Wj)

            The i-block weight's VJP — proved. Backward gradW.

            Equations
            Instances For
              noncomputable def Proofs.PairTile.tileWjHasVJP {L D C : ℕ} (Xi Xj : Mat L D) (W : Mat D C) :
              HasVJP (tileWj Xi Xj W)

              The j-block weight's VJP — proved. Backward gradWj.

              Equations
              Instances For

                Host pair planes (Layer.pairTile's pairIn) #

                With K host pair planes the layer's output is the tile with the planes P : Vec (L·(L·K)) appended — constants of the step, like the feature blocks. In the flat layout the plane block follows the tile block (finSumFinEquiv): the output is tileW … v t at inl t and P q at inr q. Each weight's Jacobian is the tile's own on the first block and zero on the second, so each VJP is the tile's backward on the cotangent's tile block (tileBlock) — the stablehlo.slice of the first C channels the pairTile backward emits when pairIn > 0.

                noncomputable def Proofs.PairTile.tileWPair {L D C K : ℕ} (Xi Xj : Mat L D) (Wj : Mat D C) (P : Vec (L * (L * K))) (v : Vec (D * C)) :
                Vec (L * (L * C) + L * (L * K))

                The tile with K constant planes appended, as a function of the i-block weight.

                Equations
                Instances For
                  noncomputable def Proofs.PairTile.tileWjPair {L D C K : ℕ} (Xi Xj : Mat L D) (W : Mat D C) (P : Vec (L * (L * K))) (v : Vec (D * C)) :
                  Vec (L * (L * C) + L * (L * K))

                  The same as a function of the j-block weight.

                  Equations
                  Instances For
                    def Proofs.PairTile.tileBlock {L C K : ℕ} (dy : Vec (L * (L * C) + L * (L * K))) :
                    Vec (L * (L * C))

                    The cotangent's tile block: its first L·(L·C) entries.

                    Equations
                    Instances For
                      theorem Proofs.PairTile.pdiv_tileWPair {L D C K : ℕ} (Xi Xj : Mat L D) (Wj : Mat D C) (P : Vec (L * (L * K))) (v : Vec (D * C)) (k : Fin (D * C)) (o : Fin (L * (L * C) + L * (L * K))) :
                      pdiv (tileWPair Xi Xj Wj P) v k o = Sum.elim (fun (t : Fin (L * (L * C))) => pdiv (tileW Xi Xj Wj) v k t) (fun (x : Fin (L * (L * K))) => 0) (finSumFinEquiv.symm o)

                      The i-block weight's Jacobian with planes appended: the tile's on the tile block, zero on the plane block.

                      theorem Proofs.PairTile.pdiv_tileWjPair {L D C K : ℕ} (Xi Xj : Mat L D) (W : Mat D C) (P : Vec (L * (L * K))) (v : Vec (D * C)) (k : Fin (D * C)) (o : Fin (L * (L * C) + L * (L * K))) :
                      pdiv (tileWjPair Xi Xj W P) v k o = Sum.elim (fun (t : Fin (L * (L * C))) => pdiv (tileWj Xi Xj W) v k t) (fun (x : Fin (L * (L * K))) => 0) (finSumFinEquiv.symm o)

                      The j-block weight's Jacobian with planes appended.

                      noncomputable def Proofs.PairTile.tileWPairHasVJP {L D C K : ℕ} (Xi Xj : Mat L D) (Wj : Mat D C) (P : Vec (L * (L * K))) :
                      HasVJP (tileWPair Xi Xj Wj P)

                      The i-block weight's VJP with planes appended — proved. Backward: gradW on the cotangent's tile block.

                      Equations
                      • One or more equations did not get rendered due to their size.
                      Instances For
                        noncomputable def Proofs.PairTile.tileWjPairHasVJP {L D C K : ℕ} (Xi Xj : Mat L D) (W : Mat D C) (P : Vec (L * (L * K))) :
                        HasVJP (tileWjPair Xi Xj W P)

                        The j-block weight's VJP with planes appended — proved. Backward: gradWj on the cotangent's tile block.

                        Equations
                        • One or more equations did not get rendered due to their size.
                        Instances For