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.
The (i, j, c) behind a flat index of a row-major [L, L, C] pair map.
Equations
- Proofs.PairTile.oidx o = ((finProdFinEquiv.symm o).1, (finProdFinEquiv.symm (finProdFinEquiv.symm o).2).1, (finProdFinEquiv.symm (finProdFinEquiv.symm o).2).2)
Instances For
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).
The pair map as a function of the flattened j-block weight (the i term is constant).
Equations
- Proofs.PairTile.tileWj Xi Xj W v o = Xi.mul W (Proofs.PairTile.oidx o).1 (Proofs.PairTile.oidx o).2.2 + Xj.mul (Proofs.Mat.unflatten v) (Proofs.PairTile.oidx o).2.1 (Proofs.PairTile.oidx o).2.2
Instances For
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
- Proofs.PairTile.gradW Xi dy k = ∑ a : Fin L, Xi a (finProdFinEquiv.symm k).1 * ∑ b : Fin L, dy (finProdFinEquiv (a, finProdFinEquiv (b, (finProdFinEquiv.symm k).2)))
Instances For
dWj (d, c) = Σ_j Xj j d · Σ_i dy i j c — summed over i, contracted with Xj over j.
Equations
- Proofs.PairTile.gradWj Xj dy k = ∑ b : Fin L, Xj b (finProdFinEquiv.symm k).1 * ∑ a : Fin L, dy (finProdFinEquiv (a, finProdFinEquiv (b, (finProdFinEquiv.symm k).2)))
Instances For
The i-block weight's VJP — proved. Backward gradW.
Equations
- Proofs.PairTile.tileWHasVJP Xi Xj Wj = { backward := fun (_v : Proofs.Vec (D * C)) (dy : Proofs.Vec (L * (L * C))) => Proofs.PairTile.gradW Xi dy, correct := ⋯ }
Instances For
The j-block weight's VJP — proved. Backward gradWj.
Equations
- Proofs.PairTile.tileWjHasVJP Xi Xj W = { backward := fun (_v : Proofs.Vec (D * C)) (dy : Proofs.Vec (L * (L * C))) => Proofs.PairTile.gradWj Xj dy, correct := ⋯ }
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.
The tile with K constant planes appended, as a function of the i-block weight.
Equations
- Proofs.PairTile.tileWPair Xi Xj Wj P v o = Sum.elim (Proofs.PairTile.tileW Xi Xj Wj v) P (finSumFinEquiv.symm o)
Instances For
The same as a function of the j-block weight.
Equations
- Proofs.PairTile.tileWjPair Xi Xj W P v o = Sum.elim (Proofs.PairTile.tileWj Xi Xj W v) P (finSumFinEquiv.symm o)
Instances For
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
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.