Imports
import Mathlib
import CLRSLean.Chapter_04.Section_04_6_Master_Theorem_All_Input4.2. Strassen’s Algorithm for Matrix Multiplication
This file formalizes Strassen's algorithm, from its algebraic 2 by 2 core all
the way to a recursive algorithm on power-of-two squares with a
Θ(n^(log₂ 7)) runtime bound.
A Matrix2 R should be read as a 2 by 2 block matrix whose entries live in
an arbitrary ring. The theorem strassen2x2_correct proves that
Strassen's seven block products reconstruct ordinary 2 by 2 block matrix
multiplication.
Recursive refinement
The type CLRS.Chapter04.SqMat is a depth-indexed square matrix:
SqMat R 0 = R and SqMat R (k+1) = Matrix (Fin 2) (Fin 2) (SqMat R k),
so SqMat R k is a genuine 2^k × 2^k matrix ring built by nesting the
2 by 2 block structure. CLRS.Chapter04.strassenRec is the recursive
seven-multiplication algorithm: it bottoms out at the scalar base case
SqMat R 0 = R and otherwise combines the seven Strassen products of its
four sub-blocks. CLRS.Chapter04.strassenRec_correct proves it computes
the ordinary matrix product A * B at every depth, and
CLRS.Chapter04.strassenRec_padOne shows that zero-padding a matrix into
the next power-of-two block preserves the product in the top-left corner.
Runtime
The work recurrence CLRS.Chapter04.strassenWork satisfies the CLRS
floor recurrence T(n) = 7 T(⌊n/2⌋) + n², i.e. seven recursive products
plus quadratic block-addition work. Feeding this into the Chapter 4 Master
theorem case-1 wrapper
CLRS.Chapter04.floorDivide_allInput_masterCase1_realLogScale gives
CLRS.Chapter04.strassen_runtime_bigTheta:
T = Θ(n^(log₂ 7)), the textbook Θ(n^(lg 7)) bound.
The companion MatrixExecution module, exported by the chapter guide,
counts all ten preparation and eight reassembly block sums/differences by
visiting their scalar entries. MatrixExecution.strassenWithCost_work_bounds
places this actual count between one and five times the budget on dimensions
2^k; MatrixExecution.strassenWithCost_theta proves its exponent.
The budget is not an exact scalar count. Padding here is a one-level embedding
of an already power-of-two square, not an arbitrary-dimension interface.
Main results:
-
Theorem
strassen2x2_correct: the 2 by 2 block algebra core. -
Theorem
strassenRec_correct: the recursive algorithm computesA * B. -
Theorem
strassenRec_padOne: zero-padding preserves the corner product. -
Theorem
strassen_runtime_bigTheta:T(n) = Θ(n^(log₂ 7)).
Notation conventions used in this section:
-
R: the scalar ring -
SqMat R k: a2^k × 2^ksquare matrix overR -
T,strassenWork: the recursive work/cost function
namespace CLRSnamespace Chapter04A 2 by 2 block matrix.
structure Matrix2 (R : Type*) where
a11 : R
a12 : R
a21 : R
a22 : Rnamespace Matrix2@[ext]
theorem ext {R : Type*} {A B : Matrix2 R}
(h11 : A.a11 = B.a11) (h12 : A.a12 = B.a12)
(h21 : A.a21 = B.a21) (h22 : A.a22 = B.a22) : A = B := by
cases A
cases B
simp_allvariable {R : Type*} [Ring R]Ordinary 2 by 2 block matrix multiplication.
def mul (A B : Matrix2 R) : Matrix2 R :=
{ a11 := A.a11 * B.a11 + A.a12 * B.a21
a12 := A.a11 * B.a12 + A.a12 * B.a22
a21 := A.a21 * B.a11 + A.a22 * B.a21
a22 := A.a21 * B.a12 + A.a22 * B.a22 }Strassen's seven-product reconstruction for 2 by 2 block matrices.
def strassen (A B : Matrix2 R) : Matrix2 R :=
let p1 := A.a11 * (B.a12 - B.a22)
let p2 := (A.a11 + A.a12) * B.a22
let p3 := (A.a21 + A.a22) * B.a11
let p4 := A.a22 * (B.a21 - B.a11)
let p5 := (A.a11 + A.a22) * (B.a11 + B.a22)
let p6 := (A.a12 - A.a22) * (B.a21 + B.a22)
let p7 := (A.a11 - A.a21) * (B.a11 + B.a12)
{ a11 := p5 + p4 - p2 + p6
a12 := p1 + p2
a21 := p3 + p4
a22 := p5 + p1 - p3 - p7 }Strassen's seven products compute the ordinary 2 by 2 block product.
theorem strassen_eq_mul (A B : Matrix2 R) : strassen A B = mul A B := by
ext <;> simp [strassen, mul] <;> noncomm_ringend Matrix2Reader-facing correctness theorem for CLRS Section 4.2: the algebraic Strassen reconstruction is extensionally equal to ordinary 2 by 2 block matrix multiplication.
theorem strassen2x2_correct {R : Type*} [Ring R] (A B : Matrix2 R) :
Matrix2.strassen A B = Matrix2.mul A B :=
Matrix2.strassen_eq_mul A B
Strassen's seven products on Matrix (Fin 2) (Fin 2)
section StrassenMatrixvariable {S : Type*} [Ring S]
Strassen's seven-product reconstruction expressed directly on
Matrix (Fin 2) (Fin 2) S. This is the Matrix-valued restatement of the
block algebra CLRS.Chapter04.Matrix2.strassen, and it is the shape used
by the recursive algorithm below. Each p is one of the seven products
P₁…P₇ from CLRS.
def strassen2 (M N : Matrix (Fin 2) (Fin 2) S) : Matrix (Fin 2) (Fin 2) S :=
let p1 := M 0 0 * (N 0 1 - N 1 1)
let p2 := (M 0 0 + M 0 1) * N 1 1
let p3 := (M 1 0 + M 1 1) * N 0 0
let p4 := M 1 1 * (N 1 0 - N 0 0)
let p5 := (M 0 0 + M 1 1) * (N 0 0 + N 1 1)
let p6 := (M 0 1 - M 1 1) * (N 1 0 + N 1 1)
let p7 := (M 0 0 - M 1 0) * (N 0 0 + N 0 1)
!![p5 + p4 - p2 + p6, p1 + p2; p3 + p4, p5 + p1 - p3 - p7]
The seven Strassen products compute the ordinary 2 × 2 matrix product. This
is the Matrix-valued counterpart of CLRS.Chapter04.Matrix2.strassen_eq_mul.
theorem strassen2_eq_mul (M N : Matrix (Fin 2) (Fin 2) S) : strassen2 M N = M * N := by
ext i j
fin_cases i <;> fin_cases j <;>
simp [strassen2, Matrix.mul_apply, Fin.sum_univ_two] <;>
noncomm_ringend StrassenMatrixRecursive Strassen on power-of-two squares
Depth-indexed square matrix over R. SqMat R 0 is the scalar type R,
and SqMat R (k+1) is a 2 × 2 block matrix whose entries are depth-k
squares. Thus SqMat R k is a 2^k × 2^k matrix realized as a balanced
quad-tree of 2 × 2 blocks.
def SqMat (R : Type u) : ℕ → Type u
| 0 => R
| (k + 1) => Matrix (Fin 2) (Fin 2) (SqMat R k)
The ring structure on CLRS.Chapter04.SqMat. At depth 0 it is the
scalar ring R; at depth k+1 it is the standard 2 × 2 matrix ring over the
depth-k ring, so ordinary multiplication on SqMat R k is exactly block
matrix multiplication.
instance instRingSqMat (R : Type u) [Ring R] : ∀ k, Ring (SqMat R k)
| 0 => inferInstanceAs (Ring R)
| (k + 1) =>
letI := instRingSqMat R k
inferInstanceAs (Ring (Matrix (Fin 2) (Fin 2) (SqMat R k)))
The recursive Strassen algorithm. strassenRec R 0 is the scalar base
case (conventional multiplication); strassenRec R (k+1) forms the seven
Strassen products of the four sub-blocks with seven recursive calls and
reassembles the four output blocks (CLRS STRASSEN).
def strassenRec (R : Type u) [Ring R] : ∀ k, SqMat R k → SqMat R k → SqMat R k
| 0, x, y => x * y
| (k + 1), A, B =>
let p1 := strassenRec R k (A 0 0) (B 0 1 - B 1 1)
let p2 := strassenRec R k (A 0 0 + A 0 1) (B 1 1)
let p3 := strassenRec R k (A 1 0 + A 1 1) (B 0 0)
let p4 := strassenRec R k (A 1 1) (B 1 0 - B 0 0)
let p5 := strassenRec R k (A 0 0 + A 1 1) (B 0 0 + B 1 1)
let p6 := strassenRec R k (A 0 1 - A 1 1) (B 1 0 + B 1 1)
let p7 := strassenRec R k (A 0 0 - A 1 0) (B 0 0 + B 0 1)
!![p5 + p4 - p2 + p6, p1 + p2; p3 + p4, p5 + p1 - p3 - p7]
Correctness of the recursive Strassen algorithm: at every depth it returns the
ordinary matrix product A * B. The proof is induction on depth; each step
rewrites the seven recursive products by the induction hypothesis and then
applies the 2 by 2 identity CLRS.Chapter04.strassen2_eq_mul.
theorem strassenRec_eq_mul (R : Type u) [Ring R] :
∀ (k : ℕ) (A B : SqMat R k), strassenRec R k A B = A * B
| 0, x, y => rfl
| (k + 1), A, B => by
have IH : ∀ X Y : SqMat R k, strassenRec R k X Y = X * Y := strassenRec_eq_mul R k
have hstep : strassenRec R (k + 1) A B = strassen2 A B := by
simp only [strassenRec, strassen2, IH]
rw [hstep]
exact strassen2_eq_mul A B
Reader-facing correctness theorem for the recursive algorithm: on a
2^k × 2^k square, CLRS.Chapter04.strassenRec produces the true matrix
product.
theorem strassenRec_correct (R : Type u) [Ring R] (k : ℕ) (A B : SqMat R k) :
strassenRec R k A B = A * B :=
strassenRec_eq_mul R k A BOne-level zero-padding of power-of-two squares
Zero-padding: embed a depth-k square into the top-left block of a depth-(k+1)
square, filling the other three blocks with zeros. This is the padding step of
CLRS STRASSEN, which enlarges an n × n input to the next power of two.
Zero-padded factors multiply block-diagonally: the top-left corner of the product of two padded matrices is the product of the two originals, with the rest still zero.
theorem padOne_mul (R : Type u) [Ring R] (k : ℕ) (x y : SqMat R k) :
padOne R k x * padOne R k y = padOne R k (x * y) := by
show (!![x, 0; 0, 0] : Matrix (Fin 2) (Fin 2) (SqMat R k)) * !![y, 0; 0, 0]
= !![x * y, 0; 0, 0]
ext i j
fin_cases i <;> fin_cases j <;>
simp [Matrix.mul_apply, Fin.sum_univ_two]
Running the recursive Strassen algorithm on two zero-padded inputs recovers the
padded product: padding to the next power of two does not change the meaningful
top-left product. This composes CLRS.Chapter04.strassenRec_correct with
CLRS.Chapter04.padOne_mul.
theorem strassenRec_padOne (R : Type u) [Ring R] (k : ℕ) (x y : SqMat R k) :
strassenRec R (k + 1) (padOne R k x) (padOne R k y) = padOne R k (x * y) := by
rw [strassenRec_correct, padOne_mul]
The top-left block projection, inverse to CLRS.Chapter04.padOne on the
padded corner. Extracting the corner after a padded Strassen multiplication
returns the original product x * y.
theorem strassenRec_padOne_corner (R : Type u) [Ring R] (k : ℕ) (x y : SqMat R k) :
(strassenRec R (k + 1) (padOne R k x) (padOne R k y)) 0 0 = x * y := by
rw [strassenRec_padOne]
show (!![x * y, 0; 0, 0] : Matrix (Fin 2) (Fin 2) (SqMat R k)) 0 0 = x * y
simp
Runtime: T(n) = 7 T(⌊n/2⌋) + n² is Θ(n^(log₂ 7))
The Strassen work recurrence T(n) = 7 T(⌊n/2⌋) + n²: seven recursive
subproblems of half size plus quadratic block-combination work, with base value
T(0) = 0. This is the CLRS cost recurrence whose solution is the running
time of CLRS.Chapter04.strassenRec.
noncomputable def strassenWork : ℕ → ℝ
| 0 => 0
| (n + 1) => 7 * strassenWork ((n + 1) / 2) + ((n + 1 : ℕ) : ℝ) ^ 2
decreasing_by exact Nat.div_lt_self (Nat.succ_pos n) (by norm_num)Base value of the work recurrence.
theorem strassenWork_zero : strassenWork 0 = 0 := by
rw [strassenWork]One recursion step of the work recurrence at a successor argument.
theorem strassenWork_succ (n : ℕ) :
strassenWork (n + 1) = 7 * strassenWork ((n + 1) / 2) + ((n + 1 : ℕ) : ℝ) ^ 2 := by
rw [strassenWork]One recursion step of the work recurrence at any positive argument.
theorem strassenWork_pos_step (n : ℕ) (hn : 0 < n) :
strassenWork n = 7 * strassenWork (n / 2) + ((n : ℕ) : ℝ) ^ 2 := by
obtain ⟨m, rfl⟩ := Nat.exists_eq_succ_of_ne_zero hn.ne'
exact strassenWork_succ m
The forcing term f(n) = T(n) - 7 T(⌊n/2⌋) of the recurrence. Choosing
f as this defect makes the CLRS floor recurrence
T(n) = 7 T(⌊n/2⌋) + f(n) hold definitionally at every input.
noncomputable def strassenForcing (n : ℕ) : ℝ :=
strassenWork n - 7 * strassenWork (n / 2)
The work function satisfies the Chapter 4 floor-division Master recurrence with
a = 7, b = 2.
theorem strassenWork_floorRec :
FloorDivideRecurrence 7 2 strassenForcing strassenWork := by
refine ⟨fun n => ?_⟩
simp only [strassenForcing]
push_cast
ringThe work function is nonnegative.
theorem strassenWork_nonneg : ∀ n, 0 ≤ strassenWork n := by
intro n
induction n using Nat.strong_induction_on with
| _ n ih =>
rcases Nat.eq_zero_or_pos n with hn | hn
· subst hn; simp [strassenWork_zero]
· rw [strassenWork_pos_step n hn]
have hlt : n / 2 < n := Nat.div_lt_self hn (by norm_num)
have hrec := ih (n / 2) hlt
nlinarith [hrec, sq_nonneg ((n : ℕ) : ℝ)]The work function is nondecreasing across one step.
theorem strassenWork_le_succ : ∀ n, strassenWork n ≤ strassenWork (n + 1) := by
intro n
induction n using Nat.strong_induction_on with
| _ n ih =>
rcases Nat.eq_zero_or_pos n with hn | hn
· subst hn; rw [strassenWork_zero]; exact strassenWork_nonneg _
· rw [strassenWork_pos_step n hn, strassenWork_pos_step (n + 1) (Nat.succ_pos n)]
have hcast : ((n : ℕ) : ℝ) ^ 2 ≤ ((n + 1 : ℕ) : ℝ) ^ 2 := by
have h1 : ((n : ℕ) : ℝ) ≤ ((n + 1 : ℕ) : ℝ) := by push_cast; linarith
have hn_nonneg : (0 : ℝ) ≤ (n : ℕ) := Nat.cast_nonneg _
nlinarith [hn_nonneg, h1]
rcases (by omega : (n + 1) / 2 = n / 2 ∨ (n + 1) / 2 = n / 2 + 1) with h | h
· rw [h]; linarith [hcast]
· rw [h]
have hj : n / 2 < n := Nat.div_lt_self hn (by norm_num)
have hstep := ih (n / 2) hj
linarith [hstep, hcast]The work function is monotone.
theorem strassenWork_monotone : Monotone strassenWork :=
monotone_nat_of_le_succ strassenWork_le_succThe work function satisfies the absolute-value monotonicity interface.
theorem strassenWork_monotoneAbs : MonotoneAbs strassenWork := by
intro m n hmn
rw [abs_of_nonneg (strassenWork_nonneg m), abs_of_nonneg (strassenWork_nonneg n)]
exact strassenWork_monotone hmn
The normalized forcing on exact powers is the convergent geometric sequence
(4/7)^(k+1): on n = 2^(k+1) the forcing is exactly the block-work
(2^(k+1))² = 4^(k+1), so dividing by 7^(k+1) gives (4/7)^(k+1).
This is what places the Strassen recurrence in Master case 1.
theorem strassen_normForcing (k : ℕ) :
normalizedForcing 7 2 strassenForcing k = (4 / 7 : ℝ) ^ (k + 1) := by
have hpos : 0 < 2 ^ (k + 1) := pow_pos (by norm_num) _
have hdiv : 2 ^ (k + 1) / 2 = 2 ^ k := by rw [pow_succ]; omega
have hcast : ((2 ^ (k + 1) : ℕ) : ℝ) ^ 2 = (4 : ℝ) ^ (k + 1) := by
push_cast
rw [show (4 : ℝ) = 2 ^ 2 by norm_num, ← pow_mul, ← pow_mul, Nat.mul_comm 2 (k + 1)]
have hforcing : strassenForcing (2 ^ (k + 1)) = (4 : ℝ) ^ (k + 1) := by
unfold strassenForcing
rw [strassenWork_pos_step (2 ^ (k + 1)) hpos, hdiv]
linarith [hcast]
unfold normalizedForcing
rw [hforcing, ← div_pow]
norm_numCase-1 hypothesis: the normalized forcing is nonnegative.
theorem strassen_term_nonneg (k : ℕ) : 0 ≤ normalizedForcing 7 2 strassenForcing k := by
rw [strassen_normForcing]; positivityCase-1 hypothesis: the normalized forcing is bounded by a geometric sequence.
theorem strassen_term_upper (k : ℕ) :
normalizedForcing 7 2 strassenForcing k ≤ (4 / 7 : ℝ) * (4 / 7 : ℝ) ^ k := by
rw [strassen_normForcing, pow_succ]
exact le_of_eq (mul_comm _ _)
Value of the work recurrence at the base input 1.
theorem strassenWork_one : strassenWork 1 = 1 := by
rw [show (1 : ℕ) = 0 + 1 from rfl, strassenWork_succ]
norm_num [strassenWork_zero]Positivity of the normalized base value, a case-1 hypothesis.
theorem strassen_base_pos : 0 < normalizedValue 7 2 strassenWork 0 := by
unfold normalizedValue
norm_num [strassenWork_one]
Runtime of Strassen's algorithm. The recurrence
T(n) = 7 T(⌊n/2⌋) + n² is Θ(n^(log₂ 7)). This is the CLRS
Θ(n^(lg 7)) bound, obtained by discharging Master-theorem case 1 (the
forcing n² is polynomially smaller than the critical n^(log₂ 7))
through the Chapter 4 wrapper
CLRS.Chapter04.floorDivide_allInput_masterCase1_realLogScale.
theorem strassen_runtime_bigTheta :
Chapter03.isBigTheta strassenWork (realLogScale 7 2) :=
floorDivide_allInput_masterCase1_realLogScale 7 2 strassenForcing strassenWork
strassenWork_floorRec (by norm_num) (by norm_num) strassenWork_monotoneAbs
strassen_base_pos strassen_term_nonneg (r := 4 / 7) (C := 4 / 7)
(by norm_num) (by norm_num) (by norm_num) strassen_term_upper
The comparison scale CLRS.Chapter04.realLogScale at a = 7,
b = 2 is the textbook power n^(log₂ 7), so
CLRS.Chapter04.strassen_runtime_bigTheta is exactly the CLRS
Θ(n^(lg 7)) statement.
theorem realLogScale_seven_two (n : ℕ) :
realLogScale 7 2 n = (n : ℝ) ^ Real.logb 2 7 := by
rw [realLogScale, realLogExponent, Real.logb]
norm_numend Chapter04end CLRSDefinitions and proofs
CLRSLean.FourthEdition.Chapter_04.MatrixExecution.Asymptotics
Asymptotic scalar work on power-of-two squares
The functions below read counters from arbitrary families of input matrices. The asymptotic variable is depth k; the comparison functions use their actual side length 2^k. Consequently these statements do not extend the input type to all natural dimensions.
namespace CLRS.Chapter04.MatrixExecution
private theorem cubic_side (k : Nat) : ((2 ^ k : Nat) : ℝ) ^ 3 = (8 : ℝ) ^ k := by
rw [Nat.cast_pow, Nat.cast_ofNat, ← pow_mul, Nat.mul_comm k 3, pow_mul]
norm_num
private theorem strassen_side (k : Nat) :
((2 ^ k : Nat) : ℝ) ^ Real.logb 2 7 = (7 : ℝ) ^ k := by
rw [Nat.cast_pow, Nat.cast_ofNat, ← Real.rpow_pow_comm (by norm_num)]
rw [Real.rpow_logb (by norm_num) (by norm_num) (by norm_num)]Every input family has cubic scalar work in its power-of-two side length.
theorem mulWithCost_theta (R : Type u) [Ring R] (A B : ∀ k, SqMat R k) :
Chapter03.isBigTheta (fun k => ((mulWithCost R k (A k) (B k)).work : ℝ))
(fun k => ((2 ^ k : Nat) : ℝ) ^ 3) := by
constructor
· rw [Chapter03.isBigO_iff]
refine ⟨2, by norm_num, 0, ?_⟩
intro k _
rw [mulWithCost_work, cubic_side, abs_of_nonneg (by positivity),
abs_of_nonneg (by positivity)]
exact_mod_cast (mulCount_bounds k).2
· rw [Chapter03.isBigOmega_iff]
refine ⟨1, by norm_num, 0, ?_⟩
intro k _
rw [mulWithCost_work, cubic_side, abs_of_nonneg (by positivity),
abs_of_nonneg (by positivity), one_mul]
exact_mod_cast (mulCount_bounds k).1Every input family has Strassen's scalar-work exponent on its side length.
theorem strassenWithCost_theta (R : Type u) [Ring R] (A B : ∀ k, SqMat R k) :
Chapter03.isBigTheta (fun k => ((strassenWithCost R k (A k) (B k)).work : ℝ))
(fun k => ((2 ^ k : Nat) : ℝ) ^ Real.logb 2 7) := by
constructor
· rw [Chapter03.isBigO_iff]
refine ⟨7, by norm_num, 0, ?_⟩
intro k _
rw [strassenWithCost_work, strassen_side, abs_of_nonneg (by positivity),
abs_of_nonneg (by positivity)]
exact_mod_cast (strassenCount_bounds k).2
· rw [Chapter03.isBigOmega_iff]
refine ⟨1, by norm_num, 0, ?_⟩
intro k _
rw [strassenWithCost_work, strassen_side, abs_of_nonneg (by positivity),
abs_of_nonneg (by positivity), one_mul]
exact_mod_cast (strassenCount_bounds k).1end CLRS.Chapter04.MatrixExecutionCLRSLean.FourthEdition.Chapter_04.MatrixExecution.Algorithms
Matrix multiplication with scalar work
Both algorithms obtain their value and work from the same recursion. Every matrix sum/difference is itself evaluated by the counted leaf traversal. Strassen's seven product results are shared; each result is computed and charged once, even when it contributes to more than one output block.
namespace CLRS.Chapter04.MatrixExecutionEight recursive products and four counted output-block sums.
def mulWithCost (R : Type u) [Ring R] :
∀ k, SqMat R k → SqMat R k → Result R k
| 0, x, y => ⟨x * y, 1⟩
| k + 1, A, B =>
let p1 := mulWithCost R k (A 0 0) (B 0 0)
let p2 := mulWithCost R k (A 0 1) (B 1 0)
let p3 := mulWithCost R k (A 0 0) (B 0 1)
let p4 := mulWithCost R k (A 0 1) (B 1 1)
let p5 := mulWithCost R k (A 1 0) (B 0 0)
let p6 := mulWithCost R k (A 1 1) (B 1 0)
let p7 := mulWithCost R k (A 1 0) (B 0 1)
let p8 := mulWithCost R k (A 1 1) (B 1 1)
let c11 := zipWithCost R (· + ·) k p1.value p2.value
let c12 := zipWithCost R (· + ·) k p3.value p4.value
let c21 := zipWithCost R (· + ·) k p5.value p6.value
let c22 := zipWithCost R (· + ·) k p7.value p8.value
⟨!![c11.value, c12.value; c21.value, c22.value],
p1.work + p2.work + p3.work + p4.work + p5.work + p6.work + p7.work + p8.work +
c11.work + c12.work + c21.work + c22.work⟩Strassen with all ten preparation and eight reassembly sums/differences counted.
def strassenWithCost (R : Type u) [Ring R] :
∀ k, SqMat R k → SqMat R k → Result R k
| 0, x, y => ⟨x * y, 1⟩
| k + 1, A, B =>
let t1 := zipWithCost R (· - ·) k (B 0 1) (B 1 1)
let t2 := zipWithCost R (· + ·) k (A 0 0) (A 0 1)
let t3 := zipWithCost R (· + ·) k (A 1 0) (A 1 1)
let t4 := zipWithCost R (· - ·) k (B 1 0) (B 0 0)
let t5 := zipWithCost R (· + ·) k (A 0 0) (A 1 1)
let t6 := zipWithCost R (· + ·) k (B 0 0) (B 1 1)
let t7 := zipWithCost R (· - ·) k (A 0 1) (A 1 1)
let t8 := zipWithCost R (· + ·) k (B 1 0) (B 1 1)
let t9 := zipWithCost R (· - ·) k (A 0 0) (A 1 0)
let t10 := zipWithCost R (· + ·) k (B 0 0) (B 0 1)
let p1 := strassenWithCost R k (A 0 0) t1.value
let p2 := strassenWithCost R k t2.value (B 1 1)
let p3 := strassenWithCost R k t3.value (B 0 0)
let p4 := strassenWithCost R k (A 1 1) t4.value
let p5 := strassenWithCost R k t5.value t6.value
let p6 := strassenWithCost R k t7.value t8.value
let p7 := strassenWithCost R k t9.value t10.value
let c11a := zipWithCost R (· + ·) k p5.value p4.value
let c11b := zipWithCost R (· - ·) k c11a.value p2.value
let c11 := zipWithCost R (· + ·) k c11b.value p6.value
let c12 := zipWithCost R (· + ·) k p1.value p2.value
let c21 := zipWithCost R (· + ·) k p3.value p4.value
let c22a := zipWithCost R (· + ·) k p5.value p1.value
let c22b := zipWithCost R (· - ·) k c22a.value p3.value
let c22 := zipWithCost R (· - ·) k c22b.value p7.value
⟨!![c11.value, c12.value; c21.value, c22.value],
t1.work + t2.work + t3.work + t4.work + t5.work + t6.work + t7.work + t8.work +
t9.work + t10.work + p1.work + p2.work + p3.work + p4.work + p5.work + p6.work +
p7.work + c11a.work + c11b.work + c11.work + c12.work + c21.work + c22a.work +
c22b.work + c22.work⟩Erasing work recovers the existing eight-product algorithm.
theorem mulWithCost_value (R : Type u) [Ring R] :
∀ (k : Nat) (A B : SqMat R k), (mulWithCost R k A B).value = mulRec R k A B := by
intro k
induction k with
| zero => intros; rfl
| succ k ih =>
intro A B
simp only [mulWithCost, addWithCost_value, ih, mulRec]Erasing work recovers the existing seven-product algorithm.
theorem strassenWithCost_value (R : Type u) [Ring R] :
∀ (k : Nat) (A B : SqMat R k),
(strassenWithCost R k A B).value = strassenRec R k A B := by
intro k
induction k with
| zero => intros; rfl
| succ k ih =>
intro A B
simp only [strassenWithCost, addWithCost_value, subWithCost_value, ih, strassenRec]Depth recurrence proved below to be the eight-product execution's work.
def mulCount : Nat → Nat
| 0 => 1
| k + 1 => 8 * mulCount k + 4 * 4 ^ kDepth recurrence including Strassen's eighteen block sums/differences.
def strassenCount : Nat → Nat
| 0 => 1
| k + 1 => 7 * strassenCount k + 18 * 4 ^ ktheorem mulWithCost_work (R : Type u) [Ring R] :
∀ (k : Nat) (A B : SqMat R k), (mulWithCost R k A B).work = mulCount k := by
intro k
induction k with
| zero => intros; rfl
| succ k ih =>
intro A B
simp only [mulWithCost, zipWithCost_work, ih, mulCount]
omegatheorem strassenWithCost_work (R : Type u) [Ring R] :
∀ (k : Nat) (A B : SqMat R k),
(strassenWithCost R k A B).work = strassenCount k := by
intro k
induction k with
| zero => intros; rfl
| succ k ih =>
intro A B
simp only [strassenWithCost, zipWithCost_work, ih, strassenCount]
omegaend CLRS.Chapter04.MatrixExecutionCLRSLean.FourthEdition.Chapter_04.MatrixExecution.Basic
Scalar-operation execution for block matrices
The counter charges one scalar addition, subtraction, or multiplication. Matrix addition/subtraction visits all four child blocks recursively; it is not charged as one scalar operation. Indexing, immutable representation, allocation, and forming the four output blocks are outside this arithmetic model. Local results are shared when a later expression uses them more than once.
namespace CLRS.Chapter04.MatrixExecutionA returned matrix and the scalar arithmetic operations used to obtain it.
structure Result (R : Type u) (k : Nat) where
value : SqMat R k
work : NatApply one scalar operation at every corresponding pair of leaves.
def zipWithCost (R : Type u) (op : R → R → R) :
∀ k, SqMat R k → SqMat R k → Result R k
| 0, x, y => ⟨op x y, 1⟩
| k + 1, A, B =>
let a := zipWithCost R op k (A 0 0) (B 0 0)
let b := zipWithCost R op k (A 0 1) (B 0 1)
let c := zipWithCost R op k (A 1 0) (B 1 0)
let d := zipWithCost R op k (A 1 1) (B 1 1)
⟨!![a.value, b.value; c.value, d.value], a.work + b.work + c.work + d.work⟩The count comes from visiting the scalar leaves of the returned matrix.
theorem zipWithCost_work (R : Type u) (op : R → R → R) :
∀ (k : Nat) (A B : SqMat R k), (zipWithCost R op k A B).work = 4 ^ k := by
intro k
induction k with
| zero => intros; rfl
| succ k ih =>
intro A B
simp only [zipWithCost, ih, pow_succ]
omegatheorem addWithCost_value (R : Type u) [Ring R] :
∀ (k : Nat) (A B : SqMat R k),
(zipWithCost R (· + ·) k A B).value = A + B := by
intro k
induction k with
| zero => intros; rfl
| succ k ih =>
intro A B
funext i j
fin_cases i <;> fin_cases j <;> simp [zipWithCost, ih] <;> rfltheorem subWithCost_value (R : Type u) [Ring R] :
∀ (k : Nat) (A B : SqMat R k),
(zipWithCost R (· - ·) k A B).value = A - B := by
intro k
induction k with
| zero => intros; rfl
| succ k ih =>
intro A B
funext i j
fin_cases i <;> fin_cases j <;> simp [zipWithCost, ih] <;> rflend CLRS.Chapter04.MatrixExecutionCLRSLean.FourthEdition.Chapter_04.MatrixExecution.Bounds
Bounds for executed scalar arithmetic
These results concern matrices of side length 2^k. The eight-product counter equals the existing work recurrence on that domain. Strassen's eighteen block sums give a different exact count, bounded by constant multiples of the existing recurrence. Neither result claims an arbitrary-dimension padding API or a machine-time bound for ring operations of nonconstant cost.
namespace CLRS.Chapter04.MatrixExecution
private theorem side_square (k : Nat) : ((2 ^ k : Nat) : ℝ) ^ 2 = (4 ^ k : Nat) := by
norm_cast
rw [← pow_mul, Nat.mul_comm k 2, pow_mul]
norm_num
theorem mulCount_eq_budget (k : Nat) : (mulCount k : ℝ) = mulWork (2 ^ k) := by
induction k with
| zero => simp [mulCount, mulWork_one]
| succ k ih =>
rw [mulCount, Nat.cast_add, Nat.cast_mul, ih,
mulWork_pos_step (2 ^ (k + 1)) (by positivity)]
have hdiv : 2 ^ (k + 1) / 2 = 2 ^ k := by rw [pow_succ]; omega
rw [hdiv, side_square]
simp only [pow_succ 4 k,
Nat.cast_mul, Nat.cast_ofNat]
norm_num
ring
theorem strassenCount_budget_bounds (k : Nat) :
strassenWork (2 ^ k) ≤ (strassenCount k : ℝ) ∧
(strassenCount k : ℝ) ≤ 5 * strassenWork (2 ^ k) := by
induction k with
| zero => norm_num [strassenCount, strassenWork_one]
| succ k ih =>
rw [strassenCount, Nat.cast_add, Nat.cast_mul,
strassenWork_pos_step (2 ^ (k + 1)) (by positivity)]
have hdiv : 2 ^ (k + 1) / 2 = 2 ^ k := by rw [pow_succ]; omega
rw [hdiv, side_square]
simp only [pow_succ 4 k,
Nat.cast_mul, Nat.cast_ofNat]
have hnonneg : (0 : ℝ) ≤ (4 ^ k : Nat) := by positivity
constructor <;> nlinarith [ih.1, ih.2]Equality with the budget is proved from the returned execution counter.
theorem mulWithCost_work_eq (R : Type u) [Ring R] (k : Nat) (A B : SqMat R k) :
((mulWithCost R k A B).work : ℝ) = mulWork (2 ^ k) := by
rw [mulWithCost_work, mulCount_eq_budget]Strassen's actual arithmetic count is within a factor of five of the budget.
theorem strassenWithCost_work_bounds (R : Type u) [Ring R]
(k : Nat) (A B : SqMat R k) :
strassenWork (2 ^ k) ≤ ((strassenWithCost R k A B).work : ℝ) ∧
((strassenWithCost R k A B).work : ℝ) ≤ 5 * strassenWork (2 ^ k) := by
rw [strassenWithCost_work]
exact strassenCount_budget_bounds ktheorem mulCount_closed (k : Nat) : mulCount k + 4 ^ k = 2 * 8 ^ k := by
induction k with
| zero => norm_num [mulCount]
| succ k ih => simp only [mulCount, pow_succ]; omegatheorem strassenCount_closed (k : Nat) : strassenCount k + 6 * 4 ^ k = 7 * 7 ^ k := by
induction k with
| zero => norm_num [strassenCount]
| succ k ih => simp only [strassenCount, pow_succ]; omega
theorem mulCount_bounds (k : Nat) : 8 ^ k ≤ mulCount k ∧ mulCount k ≤ 2 * 8 ^ k := by
have h := mulCount_closed k
have hpow : 4 ^ k ≤ 8 ^ k := Nat.pow_le_pow_left (by norm_num) k
have hnonneg := Nat.zero_le (4 ^ k)
omega
theorem strassenCount_bounds (k : Nat) :
7 ^ k ≤ strassenCount k ∧ strassenCount k ≤ 7 * 7 ^ k := by
have h := strassenCount_closed k
have hpow : 4 ^ k ≤ 7 ^ k := Nat.pow_le_pow_left (by norm_num) k
omegaValue correctness and work for the same eight-product run.
theorem mulWithCost_correct (R : Type u) [Ring R] (k : Nat) (A B : SqMat R k) :
(mulWithCost R k A B).value = A * B ∧
((mulWithCost R k A B).work : ℝ) = mulWork (2 ^ k) :=
⟨by rw [mulWithCost_value, mulRec_correct], mulWithCost_work_eq R k A B⟩Value correctness and work bounds for the same seven-product run.
theorem strassenWithCost_correct (R : Type u) [Ring R] (k : Nat) (A B : SqMat R k) :
(strassenWithCost R k A B).value = A * B ∧
strassenWork (2 ^ k) ≤ ((strassenWithCost R k A B).work : ℝ) ∧
((strassenWithCost R k A B).work : ℝ) ≤ 5 * strassenWork (2 ^ k) :=
⟨by rw [strassenWithCost_value, strassenRec_correct], strassenWithCost_work_bounds R k A B⟩end CLRS.Chapter04.MatrixExecutionCLRSLean.Chapter_04.Section_04_6_Master_Theorem_All_Input
The real-log comparison scale n^(log_b a). This is the textbook scale used in
the standard CLRS statement of the Master theorem: the homogeneous-solution
growth rate without floors and ceilings.
For integer exponents it coincides with the ordinary polynomial scale
polynomialScale.
noncomputable def realLogScale (a b : ℕ) (n : ℕ) : ℝ :=
(n : ℝ) ^ (realLogExponent a b)
Floor-division all-input Master case 1 stated in the textbook real-log scale
n^(log_b a).
theorem floorDivide_allInput_masterCase1_realLogScale
(a b : ℕ) (f T : ℕ → ℝ)
(h_rec : FloorDivideRecurrence a b f T)
(ha : 1 ≤ a) (hb : 1 < b)
(hT_mono : MonotoneAbs T)
(h_base_pos : 0 < normalizedValue a b T 0)
(h_term_nonneg : ∀ k, 0 ≤ normalizedForcing a b f k)
{r C : ℝ} (hr_nonneg : 0 ≤ r) (hr_lt_one : r < 1) (hC_pos : 0 < C)
(h_term_upper : ∀ k, normalizedForcing a b f k ≤ C * r ^ k) :
Chapter03.isBigTheta T (realLogScale a b) := by
exact Chapter03.isBigTheta_trans
(floorDivide_allInput_masterCase1_criticalPowerScale a b f T h_rec
ha hb hT_mono h_base_pos h_term_nonneg hr_nonneg hr_lt_one hC_pos
h_term_upper)
(criticalPowerScale_isBigTheta_realLogScale a b ha hb)