Documentation

LeanMlir.Proofs.Foundation.Bf16Erasure

Erasure — every batched bf16 kind is its f32 peer at the identity rounding #

flatConvFBf16_id (StableHLO.Basic) says of the per-example CIFAR conv that the bf16 op "adds ROUNDING and nothing else — no reassociation, no dropped bias, no moved padding". This file says it of the twenty-five batched kinds the ImageNet renders emit: for each, den (or denOp) at rnd := id is the f32 constructor's den, named against its own peer. The renderers pass zrnd = id (StableHLO.Pretty), so these are the equalities the rendered bf16 ASTs satisfy, and they are what lets a whole-net statement at the f32 nodes be restated at the bf16 artifact (planning/bf16_tie.md §3).

Each lemma names its own peer. The symmetric and XLA-SAME strided kinds have identical types and emitted shapes and differ only in denOp (Basic.lean, the convStridedXla arm), so a mismatched pairing fails to prove — which is the check.

groupkindsproof
forward conv / depthwise (7)convBf16, convStridedBf16, convStridedXlaBf16, convStride4Bf16, depthwiseBf16, depthwiseStridedBf16, depthwiseStridedXlaBf16the bias sits outside the store in the bf16 den and inside the conv in the f32 one: *_bias_split
forward dense / patch (2)denseRowBf16, patchEmbedBf16rowBiasFlat after the store vs dense's bias; patchEmbedFlatBf16 id is patchEmbedFlat
denseRowBackBf16, matmulFBBf16rfl
dgrad (5)convBackBatchedBf16, convStridedBackBatchedBf16, depthwiseBackBatchedBf16, depthwiseStridedBackBatchedBf16, depthwiseStridedXlaBackBatchedBf16rfl
wgrad (9)the Bf16GradNodes tablerfl

Two bias splits the suite did not have (flatConvStride2, depthwiseStride2FlatXla) are proved here beside their siblings' pattern from ParamGradNodes.

The last section states the same thing of the renderers' switches (StableHLO.PrecisionSwitch): denOp (.convAt bf16 id …) = denOp (.conv …) for either bf16, one lemma per switch. A typed forward graph built on the switches (r34IdGraphB, …) is then faithful at either precision by the f32 proof with these rewrites in front of it — simp only […, denOp_convAt_id] before denOp, since on a symbolic bf16 the denOp equations are stuck on the if and simp would otherwise unfold it to a stuck match.

Nothing here says how large the rounding is; rnd is a binder everywhere else and id here.

theorem Proofs.Bf16Fold.flatConvStride2_bias_split {ic oc h w kH kW : ℕ} (W : Kernel4 oc ic kH kW) (θ : Vec oc) (y : Vec (ic * (2 * h) * (2 * w))) :
flatConvStride2 W θ y = fun (k : Fin (oc * h * w)) => flatConvStride2 W 0 y k + broadcastFlat oc h w θ k

A symmetric strided conv's bias is a channel broadcast: decimation keeps channels.

theorem Proofs.Bf16Fold.depthwiseStride2FlatXla_bias_split {c h w kH kW : ℕ} (W : DepthwiseKernel c kH kW) (θ : Vec c) (y : Vec (c * (2 * h) * (2 * w))) :
depthwiseStride2FlatXla W θ y = fun (k : Fin (c * h * w)) => depthwiseStride2FlatXla W 0 y k + broadcastFlat c h w θ k

An XLA-SAME strided depthwise conv's bias is a channel broadcast.

.batchOp lifts a denOp equality: the batched node at either precision.

theorem Proofs.Bf16Fold.convBf16_id {ic oc h w kH kW : ℕ} (wN bN : String) (W : Kernel4 oc ic kH kW) (bias : Vec oc) :
theorem Proofs.Bf16Fold.convStridedBf16_id {ic oc h w kH kW : ℕ} (wN bN : String) (W : Kernel4 oc ic kH kW) (bias : Vec oc) :
theorem Proofs.Bf16Fold.convStride4Bf16_id {ic oc h w kH kW : ℕ} (wN bN : String) (W : Kernel4 oc ic kH kW) (bias : Vec oc) :

The row-dense forward: the bf16 store sits before the bias (rowBiasFlat after it), the f32 dense carries its bias inside; at id the two are the same affine map.

theorem Proofs.Bf16Fold.patchEmbedBf16_id {ic H W P N D : ℕ} (wN bN cN pN : String) (Wc : Kernel4 D ic P P) (bc cls : Vec D) (pos : Mat (N + 1) D) :
theorem Proofs.Bf16Fold.convBackBatchedBf16_id {N ic oc h w kH kW : ℕ} (wN : String) (W : Kernel4 oc ic kH kW) (b : Vec oc) (e : StableHLO.SHlo (N * (oc * h * w))) :
theorem Proofs.Bf16Fold.convWeightGradBBf16_id {N ic oc h w kH kW : ℕ} (xN : String) (b : Vec oc) (x : Vec (N * (ic * h * w))) (W : Kernel4 oc ic kH kW) (e : StableHLO.SHlo (N * (oc * h * w))) :
theorem Proofs.Bf16Fold.convStridedWeightGradBBf16_id {N ic oc h w kH kW : ℕ} (xN : String) (b : Vec oc) (x : Vec (N * (ic * (2 * h) * (2 * w)))) (W : Kernel4 oc ic kH kW) (e : StableHLO.SHlo (N * (oc * h * w))) :
theorem Proofs.Bf16Fold.convStridedXlaWeightGradBBf16_id {N ic oc h w kH kW : ℕ} (xN : String) (b : Vec oc) (x : Vec (N * (ic * (2 * h) * (2 * w)))) (W : Kernel4 oc ic kH kW) (e : StableHLO.SHlo (N * (oc * h * w))) :
theorem Proofs.Bf16Fold.convStride4WeightGradBBf16_id {N ic oc h w kH kW : ℕ} (xN : String) (b : Vec oc) (x : Vec (N * (ic * (2 * (2 * h)) * (2 * (2 * w))))) (W : Kernel4 oc ic kH kW) (e : StableHLO.SHlo (N * (oc * h * w))) :
theorem Proofs.Bf16Fold.depthwiseWeightGradBBf16_id {N c h w kH kW : ℕ} (xN : String) (b : Vec c) (x : Vec (N * (c * h * w))) (W : DepthwiseKernel c kH kW) (e : StableHLO.SHlo (N * (c * h * w))) :
theorem Proofs.Bf16Fold.denOp_convAt_id (bf16 : Bool) {ic oc h w kH kW : ℕ} (wN bN : String) (W : Kernel4 oc ic kH kW) (b : Vec oc) :
theorem Proofs.Bf16Fold.denOp_convStridedAt_id (bf16 : Bool) {ic oc h w kH kW : ℕ} (wN bN : String) (W : Kernel4 oc ic kH kW) (b : Vec oc) :
theorem Proofs.Bf16Fold.denOp_convStride4At_id (bf16 : Bool) {ic oc h w kH kW : ℕ} (wN bN : String) (W : Kernel4 oc ic kH kW) (b : Vec oc) :
theorem Proofs.Bf16Fold.denOp_patchEmbedAt_id (bf16 : Bool) {ic H W P N D : ℕ} (wN bN cN pN : String) (Wc : Kernel4 D ic P P) (bc cls : Vec D) (pos : Mat (N + 1) D) :
StableHLO.denOp (StableHLO.BatchableOp.patchEmbedAt bf16 id wN bN cN pN Wc bc cls pos) = StableHLO.denOp (StableHLO.BatchableOp.patchEmbed wN bN cN pN Wc bc cls pos)
theorem Proofs.Bf16Fold.den_flatConvFAt_id (bf16 : Bool) {ic oc h w kH kW : ℕ} (wN bN : String) (W : Kernel4 oc ic kH kW) (b : Vec oc) (e : StableHLO.SHlo (ic * h * w)) :
theorem Proofs.Bf16Fold.den_convBackBatchedAt_id (bf16 : Bool) {N ic oc h w kH kW : ℕ} (wN : String) (W : Kernel4 oc ic kH kW) (b : Vec oc) (e : StableHLO.SHlo (N * (oc * h * w))) :
theorem Proofs.Bf16Fold.den_convStridedBackBatchedAt_id (bf16 : Bool) {N ic oc h w kH kW : ℕ} (wN : String) (W : Kernel4 oc ic kH kW) (b : Vec oc) (e : StableHLO.SHlo (N * (oc * h * w))) :
theorem Proofs.Bf16Fold.den_convWeightGradBAt_id (bf16 : Bool) {N ic oc h w kH kW : ℕ} (xN : String) (b : Vec oc) (x : Vec (N * (ic * h * w))) (W : Kernel4 oc ic kH kW) (e : StableHLO.SHlo (N * (oc * h * w))) :
theorem Proofs.Bf16Fold.den_convStridedWeightGradBAt_id (bf16 : Bool) {N ic oc h w kH kW : ℕ} (xN : String) (b : Vec oc) (x : Vec (N * (ic * (2 * h) * (2 * w)))) (W : Kernel4 oc ic kH kW) (e : StableHLO.SHlo (N * (oc * h * w))) :
theorem Proofs.Bf16Fold.den_convStridedXlaWeightGradBAt_id (bf16 : Bool) {N ic oc h w kH kW : ℕ} (xN : String) (b : Vec oc) (x : Vec (N * (ic * (2 * h) * (2 * w)))) (W : Kernel4 oc ic kH kW) (e : StableHLO.SHlo (N * (oc * h * w))) :
theorem Proofs.Bf16Fold.den_convStride4WeightGradBAt_id (bf16 : Bool) {N ic oc h w kH kW : ℕ} (xN : String) (b : Vec oc) (x : Vec (N * (ic * (2 * (2 * h)) * (2 * (2 * w))))) (W : Kernel4 oc ic kH kW) (e : StableHLO.SHlo (N * (oc * h * w))) :
theorem Proofs.Bf16Fold.den_depthwiseWeightGradBAt_id (bf16 : Bool) {N c h w kH kW : ℕ} (xN : String) (b : Vec c) (x : Vec (N * (c * h * w))) (W : DepthwiseKernel c kH kW) (e : StableHLO.SHlo (N * (c * h * w))) :
theorem Proofs.Bf16Fold.den_depthwiseStridedWeightGradBAt_id (bf16 : Bool) {N c h w kH kW : ℕ} (xN : String) (b : Vec c) (x : Vec (N * (c * (2 * h) * (2 * w)))) (W : DepthwiseKernel c kH kW) (e : StableHLO.SHlo (N * (c * h * w))) :
theorem Proofs.Bf16Fold.den_depthwiseStridedXlaWeightGradBAt_id (bf16 : Bool) {N c h w kH kW : ℕ} (xN : String) (b : Vec c) (x : Vec (N * (c * (2 * h) * (2 * w)))) (W : DepthwiseKernel c kH kW) (e : StableHLO.SHlo (N * (c * h * w))) :