Chapter 30 — Polynomials and the FFT
CLRS, fourth edition · Lean 4 formalization
The proofs below use the models and assumptions described in the scope and implementation notes.
Imports
import CLRSLean.FourthEdition.Chapter_30.Section_30_1_Representing_Polynomials.S1_CoefficientVectors
import CLRSLean.FourthEdition.Chapter_30.Section_30_1_Representing_Polynomials.S2_PointValueInterpolation
import CLRSLean.FourthEdition.Chapter_30.Section_30_1_Representing_Polynomials.S3_RepresentationOperations30.1. Representing Polynomials
Fixed-capacity coefficient vectors are connected to Polynomial by exact
round trips, and hornerEval_correct verifies their canonical Horner
execution. Distinct point-value samples determine bounded-degree
polynomials through interpolate_pointValues_roundTrip. Vector addition,
pointwise multiplication, and the explicit pair-traversing schoolbook
execution have representation theorems and execution-attached exact costs.
Roots of unity and Fourier transforms belong to Section 30.2.
Implementation pages:
namespace CLRSnamespace Chapter30end Chapter30end CLRSDefinitions and proofs
CLRSLean.FourthEdition.Chapter_30.Section_30_1_Representing_Polynomials.S1_CoefficientVectors
Chapter 30.1: Fixed-capacity coefficient vectors
This module establishes the representation boundary between mathematical polynomials and the total fixed-capacity vectors used by Chapter 30's algorithms.
namespace CLRSnamespace Chapter30open Polynomial
A total coefficient vector with fixed capacity n.
abbrev CoeffVector (K : Type*) (n : Nat) := Fin n → KA coefficient vector whose capacity is a power of two.
abbrev PowTwoVec (K : Type*) (k : Nat) := CoeffVector K (2 ^ k)
Read and zero-pad the first n coefficients of a polynomial.
def coeffVector [Semiring K] (n : Nat) (p : K[X]) : CoeffVector K n :=
fun i => p.coeff iReconstruct a polynomial from every slot of a fixed coefficient vector.
noncomputable def vectorToPolynomial [Semiring K] {n : Nat}
(a : CoeffVector K n) : K[X] :=
∑ i : Fin n, Polynomial.monomial i.1 (a i)Reconstruction preserves every coefficient whose index is in range.
theorem vectorToPolynomial_coeff [Semiring K] {n : Nat}
(a : CoeffVector K n) (i : Fin n) :
(vectorToPolynomial a).coeff i = a i := by
classical
change (Polynomial.lcoeff K i)
(∑ j : Fin n, Polynomial.monomial j.1 (a j)) = a i
rw [map_sum]
simp only [Polynomial.lcoeff_apply, Polynomial.coeff_monomial]
rw [Finset.sum_eq_single i]
· simp
· intro j _ hji
rw [if_neg]
exact fun hval => hji (Fin.ext hval)
· simpReading the coefficients of a reconstructed vector returns that vector.
theorem coeffVector_vectorToPolynomial [Semiring K] {n : Nat}
(a : CoeffVector K n) :
coeffVector n (vectorToPolynomial a) = a := by
funext i
exact vectorToPolynomial_coeff a iA reconstructed vector has no coefficient at or beyond its capacity.
theorem vectorToPolynomial_coeff_eq_zero_of_ge [Semiring K] {n i : Nat}
(a : CoeffVector K n) (hi : n ≤ i) :
(vectorToPolynomial a).coeff i = 0 := by
classical
have hne : ∀ j : Fin n, (j : Nat) ≠ i := by
intro j hji
omega
simp [vectorToPolynomial, Polynomial.coeff_monomial, hne]The degree of a reconstructed vector is strictly below its capacity.
theorem vectorToPolynomial_degree_lt [Semiring K] {n : Nat}
(a : CoeffVector K n) :
(vectorToPolynomial a).degree < n := by
rw [Polynomial.degree_lt_iff_coeff_zero]
intro i hi
exact vectorToPolynomial_coeff_eq_zero_of_ge a hiReconstruction after truncation is exact when the polynomial fits.
theorem vectorToPolynomial_coeffVector [Semiring K] {n : Nat} (p : K[X])
(hp : p.degree < n) :
vectorToPolynomial (coeffVector n p) = p := by
ext i
by_cases hi : i < n
· let j : Fin n := ⟨i, hi⟩
simpa [coeffVector, j] using
(vectorToPolynomial_coeff (coeffVector n p) j)
· rw [vectorToPolynomial_coeff_eq_zero_of_ge]
· exact ((Polynomial.degree_lt_iff_coeff_zero p n).mp hp) i
(Nat.le_of_not_gt hi) |>.symm
· exact Nat.le_of_not_gt hiRemove the constant slot from a nonempty low-coefficient-first vector.
def tailCoeffs {K : Type*} {n : Nat} (a : CoeffVector K (n + 1)) :
CoeffVector K n :=
fun i => a i.succThe value and arithmetic counters produced by a scalar computation.
The computed scalar.
The number of charged additions.
The number of charged multiplications.
structure ArithmeticExecution (K : Type*) where value : K additions : Nat multiplications : NatTotal charged arithmetic work of a scalar execution.
def ArithmeticExecution.work (r : ArithmeticExecution K) : Nat :=
r.additions + r.multiplicationsCanonical Horner execution on a low-coefficient-first vector.
def hornerEvalExec [Semiring K] :
{n : Nat} → CoeffVector K n → K → ArithmeticExecution K
| 0, _, _ => ⟨0, 0, 0⟩
| n + 1, a, x =>
let child := hornerEvalExec (tailCoeffs a) x
⟨a 0 + x * child.value,
child.additions + 1,
child.multiplications + 1⟩The value returned by the canonical Horner execution.
def hornerEval [Semiring K] {n : Nat} (a : CoeffVector K n) (x : K) : K :=
(hornerEvalExec a x).valueReconstructing a nonempty vector splits into its constant term and shifted tail polynomial.
theorem vectorToPolynomial_succ [CommSemiring K] {n : Nat}
(a : CoeffVector K (n + 1)) :
vectorToPolynomial a = Polynomial.monomial 0 (a 0) +
Polynomial.X * vectorToPolynomial (tailCoeffs a) := by
simp [vectorToPolynomial, tailCoeffs, Fin.sum_univ_succ, Finset.mul_sum,
Polynomial.X, Polynomial.monomial_mul_monomial, Nat.add_comm]Evaluation of a nonempty coefficient vector splits into its constant term and the shifted tail.
theorem vectorToPolynomial_eval_succ [CommSemiring K] {n : Nat}
(a : CoeffVector K (n + 1)) (x : K) :
(vectorToPolynomial a).eval x =
a 0 + x * (vectorToPolynomial (tailCoeffs a)).eval x := by
rw [vectorToPolynomial_succ]
simpHorner execution evaluates the polynomial represented by its input.
theorem hornerEval_correct [CommSemiring K] {n : Nat}
(a : CoeffVector K n) (x : K) :
hornerEval a x = (vectorToPolynomial a).eval x := by
induction n with
| zero => simp [hornerEval, hornerEvalExec, vectorToPolynomial]
| succ n ih =>
simp only [hornerEval, hornerEvalExec]
change a 0 + x * hornerEval (tailCoeffs a) x =
(vectorToPolynomial a).eval x
rw [ih (tailCoeffs a)]
exact (vectorToPolynomial_eval_succ a x).symmHorner execution charges exactly one addition per input slot.
theorem hornerEvalExec_additions [Semiring K] {n : Nat}
(a : CoeffVector K n) (x : K) :
(hornerEvalExec a x).additions = n := by
induction n with
| zero => rfl
| succ n ih =>
simp only [hornerEvalExec]
rw [ih (tailCoeffs a)]Horner execution charges exactly one multiplication per input slot.
theorem hornerEvalExec_multiplications [Semiring K] {n : Nat}
(a : CoeffVector K n) (x : K) :
(hornerEvalExec a x).multiplications = n := by
induction n with
| zero => rfl
| succ n ih =>
simp only [hornerEvalExec]
rw [ih (tailCoeffs a)]Horner execution performs exactly twice the vector capacity in charged arithmetic operations.
theorem hornerEvalWork_exact [Semiring K] {n : Nat}
(a : CoeffVector K n) (x : K) :
(hornerEvalExec a x).work = 2 * n := by
rw [ArithmeticExecution.work, hornerEvalExec_additions,
hornerEvalExec_multiplications]
omegaend Chapter30end CLRSCLRSLean.FourthEdition.Chapter_30.Section_30_1_Representing_Polynomials.S2_PointValueInterpolation
Chapter 30.1: Point-value interpolation
Point-value vectors are connected to Mathlib's Lagrange interpolant. The public theorems keep the distinct-node and degree-capacity premises explicit.
namespace CLRSnamespace Chapter30open PolynomialEvaluate a polynomial at every node in a fixed vector.
def pointValues [Semiring K] {n : Nat} (points : Fin n → K) (p : K[X]) :
CoeffVector K n :=
fun i => p.eval (points i)The Lagrange interpolant through a fixed vector of nodes and values.
noncomputable def interpolateVector [Field K] {n : Nat}
(points values : Fin n → K) : K[X] :=
Lagrange.interpolate Finset.univ points valuesTwo degree-bounded polynomials agreeing at distinct nodes are equal.
theorem pointValues_injective [Field K] {n : Nat}
{points : Fin n → K} (hpoints : Function.Injective points)
{p q : K[X]} (hp : p.degree < n) (hq : q.degree < n)
(hvalues : pointValues points p = pointValues points q) : p = q := by
apply Polynomial.eq_of_degrees_lt_of_eval_index_eq Finset.univ
hpoints.injOn
· simpa using hp
· simpa using hq
· intro i _
exact congrFun hvalues iThe Lagrange interpolant assumes its prescribed value at every distinct sample node.
theorem interpolate_pointValues [Field K] {n : Nat}
{points values : Fin n → K} (hpoints : Function.Injective points)
(i : Fin n) :
(interpolateVector points values).eval (points i) = values i := by
exact Lagrange.eval_interpolate_at_node values hpoints.injOn (by simp)Every Lagrange basis divisor has natural degree at most one.
private theorem basisDivisor_natDegree_le_one [Field K] (x y : K) :
(Lagrange.basisDivisor x y).natDegree ≤ 1 := by
by_cases hxy : x = y
· simp [hxy, Lagrange.basisDivisor_self]
· simp [Lagrange.natDegree_basisDivisor_of_ne hxy]A Lagrange basis polynomial has natural degree at most the number of factors in its defining product.
private theorem lagrangeBasis_natDegree_le [Field K] {ι : Type*}
[DecidableEq ι] (s : Finset ι) (points : ι → K) (i : ι) :
(Lagrange.basis s points i).natDegree ≤ (s.erase i).card := by
rw [Lagrange.basis]
calc
(∏ j ∈ s.erase i, Lagrange.basisDivisor (points i) (points j)).natDegree
≤ ∑ j ∈ s.erase i,
(Lagrange.basisDivisor (points i) (points j)).natDegree :=
Polynomial.natDegree_prod_le _ _
_ ≤ ∑ _j ∈ s.erase i, 1 := by
exact Finset.sum_le_sum fun j _ =>
basisDivisor_natDegree_le_one (points i) (points j)
_ = (s.erase i).card := by simpThe interpolant always fits in the declared node capacity, independently of whether the nodes are distinct.
theorem interpolateVector_degree_lt [Field K] {n : Nat}
(points values : Fin n → K) :
(interpolateVector points values).degree < n := by
cases n with
| zero =>
simp [interpolateVector, Lagrange.interpolate_empty]
| succ n =>
have hnat : (interpolateVector points values).natDegree ≤ n := by
rw [interpolateVector, Lagrange.interpolate_apply]
apply Polynomial.natDegree_sum_le_of_forall_le
intro i hi
calc
(Polynomial.C (values i) *
Lagrange.basis Finset.univ points i).natDegree
≤ (Polynomial.C (values i)).natDegree +
(Lagrange.basis Finset.univ points i).natDegree :=
Polynomial.natDegree_mul_le
_ ≤ 0 + (Finset.univ.erase i).card := by
gcongr
· simp
· exact lagrangeBasis_natDegree_le Finset.univ points i
_ = n := by
rw [Finset.card_erase_of_mem hi]
simp
calc
(interpolateVector points values).degree ≤
((interpolateVector points values).natDegree : WithBot Nat) :=
Polynomial.degree_le_natDegree
_ ≤ (n : WithBot Nat) := by exact_mod_cast hnat
_ < ((n + 1 : Nat) : WithBot Nat) := by exact_mod_cast Nat.lt_succ_self nA degree-bounded polynomial with prescribed values is the interpolant.
theorem interpolate_unique [Field K] {n : Nat}
{points values : Fin n → K} (hpoints : Function.Injective points)
{p : K[X]} (hp : p.degree < n)
(heval : ∀ i, p.eval (points i) = values i) :
p = interpolateVector points values := by
apply pointValues_injective hpoints hp
(interpolateVector_degree_lt points values)
funext i
rw [pointValues, pointValues, heval i, interpolate_pointValues hpoints i]Interpolating the point values of a fitting polynomial is a round trip.
theorem interpolate_pointValues_roundTrip [Field K] {n : Nat}
{points : Fin n → K} (hpoints : Function.Injective points)
{p : K[X]} (hp : p.degree < n) :
interpolateVector points (pointValues points p) = p := by
exact (interpolate_unique hpoints hp (fun _ => rfl)).symmend Chapter30end CLRSCLRSLean.FourthEdition.Chapter_30.Section_30_1_Representing_Polynomials.S3_RepresentationOperations
Chapter 30.1: Representation operations
This module gives canonical executions for fixed-vector operations. Its schoolbook multiplication really traverses every coefficient pair, inserts the product into one output bucket, and increments the attached counters.
namespace CLRSnamespace Chapter30open PolynomialThe value and arithmetic counters produced by a fixed-vector operation.
The computed vector.
The number of charged additions.
The number of charged multiplications.
structure VectorArithmeticExecution (K : Type*) (n : Nat) where value : CoeffVector K n additions : Nat multiplications : NatTotal charged arithmetic work of a vector execution.
def VectorArithmeticExecution.work (r : VectorArithmeticExecution K n) : Nat :=
r.additions + r.multiplicationsCanonical pointwise vector-addition execution.
def vectorAddExec [AddMonoid K] {n : Nat}
(a b : CoeffVector K n) : VectorArithmeticExecution K n :=
⟨fun i => a i + b i, n, 0⟩The value returned by canonical vector addition.
def vectorAdd [AddMonoid K] {n : Nat}
(a b : CoeffVector K n) : CoeffVector K n :=
(vectorAddExec a b).valueCanonical pointwise vector-multiplication execution.
def pointwiseMulExec [Mul K] {n : Nat}
(a b : CoeffVector K n) : VectorArithmeticExecution K n :=
⟨fun i => a i * b i, 0, n⟩The value returned by canonical pointwise multiplication.
def pointwiseMul [Mul K] {n : Nat}
(a b : CoeffVector K n) : CoeffVector K n :=
(pointwiseMulExec a b).valueVector addition charges exactly one addition per slot.
theorem vectorAddWork_exact [AddMonoid K] {n : Nat}
(a b : CoeffVector K n) :
(vectorAddExec a b).work = n := by
simp [vectorAddExec, VectorArithmeticExecution.work]Pointwise multiplication charges exactly one multiplication per slot.
theorem pointwiseMulWork_exact [Mul K] {n : Nat}
(a b : CoeffVector K n) :
(pointwiseMulExec a b).work = n := by
simp [pointwiseMulExec, VectorArithmeticExecution.work]Vector addition represents polynomial addition.
theorem vectorToPolynomial_vectorAdd [CommSemiring K] {n : Nat}
(a b : CoeffVector K n) :
vectorToPolynomial (vectorAdd a b) =
vectorToPolynomial a + vectorToPolynomial b := by
ext i
by_cases hi : i < n
· let j : Fin n := ⟨i, hi⟩
change (vectorToPolynomial (vectorAdd a b)).coeff j =
(vectorToPolynomial a + vectorToPolynomial b).coeff j
rw [Polynomial.coeff_add, vectorToPolynomial_coeff,
vectorToPolynomial_coeff, vectorToPolynomial_coeff]
rfl
· have hge : n ≤ i := Nat.le_of_not_gt hi
simp [vectorToPolynomial_coeff_eq_zero_of_ge, hge]Sampling commutes with polynomial addition.
theorem pointValues_add [Semiring K] {n : Nat} (points : Fin n → K)
(p q : K[X]) :
pointValues points (p + q) =
vectorAdd (pointValues points p) (pointValues points q) := by
funext i
simp [pointValues, vectorAdd, vectorAddExec]Sampling commutes with polynomial multiplication.
theorem pointValues_mul [CommSemiring K] {n : Nat} (points : Fin n → K)
(p q : K[X]) :
pointValues points (p * q) =
pointwiseMul (pointValues points p) (pointValues points q) := by
funext i
simp [pointValues, pointwiseMul, pointwiseMulExec]
The output bucket receiving the product of input slots i and j.
def productIndex {m n : Nat} (i : Fin m) (j : Fin n) : Fin (m + n) :=
⟨i.1 + j.1, by omega⟩Add a scalar to one output bucket and leave every other bucket unchanged.
def addToBucket [AddMonoid K] {n : Nat} (out : CoeffVector K n)
(i : Fin n) (v : K) : CoeffVector K n :=
Function.update out i (out i + v)A vector supported at exactly one output bucket.
def singleBucket [Zero K] {n : Nat} (i : Fin n) (v : K) : CoeffVector K n :=
fun j => if j = i then v else 0Updating one bucket is pointwise addition by its singleton vector.
theorem addToBucket_eq_vectorAdd_singleBucket [AddMonoid K] {n : Nat}
(out : CoeffVector K n) (i : Fin n) (v : K) :
addToBucket out i v = vectorAdd out (singleBucket i v) := by
funext j
by_cases hji : j = i
· subst j
simp [addToBucket, vectorAdd, vectorAddExec, singleBucket]
· simp [addToBucket, vectorAdd, vectorAddExec, singleBucket, hji]Reconstructing a singleton bucket gives the corresponding monomial.
theorem vectorToPolynomial_singleBucket [Semiring K] {n : Nat}
(i : Fin n) (v : K) :
vectorToPolynomial (singleBucket i v) = Polynomial.monomial i.1 v := by
ext k
by_cases hk : k < n
· let j : Fin n := ⟨k, hk⟩
simpa [singleBucket, j, Polynomial.coeff_monomial, Fin.ext_iff,
eq_comm] using (vectorToPolynomial_coeff (singleBucket i v) j)
· have hge : n ≤ k := Nat.le_of_not_gt hk
have hik : (i : Nat) ≠ k := by omega
rw [vectorToPolynomial_coeff_eq_zero_of_ge _ hge]
simp [Polynomial.coeff_monomial, hik]One bucket insertion adds exactly its corresponding monomial.
theorem vectorToPolynomial_addToBucket [CommSemiring K] {n : Nat}
(out : CoeffVector K n) (i : Fin n) (v : K) :
vectorToPolynomial (addToBucket out i v) =
vectorToPolynomial out + Polynomial.monomial i.1 v := by
rw [addToBucket_eq_vectorAdd_singleBucket,
vectorToPolynomial_vectorAdd, vectorToPolynomial_singleBucket]One coefficient-pair step of the schoolbook execution.
def schoolbookStep [Semiring K] {m n : Nat}
(a : CoeffVector K m) (b : CoeffVector K n)
(r : VectorArithmeticExecution K (m + n)) (ij : Fin m × Fin n) :
VectorArithmeticExecution K (m + n) :=
⟨addToBucket r.value (productIndex ij.1 ij.2) (a ij.1 * b ij.2),
r.additions + 1,
r.multiplications + 1⟩Traverse a list of coefficient pairs from an arbitrary execution state.
def schoolbookPairsExec [Semiring K] {m n : Nat}
(a : CoeffVector K m) (b : CoeffVector K n)
(pairs : List (Fin m × Fin n))
(initial : VectorArithmeticExecution K (m + n)) :
VectorArithmeticExecution K (m + n) :=
pairs.foldl (schoolbookStep a b) initialThe zero-valued initial state of schoolbook multiplication.
def schoolbookInitial [Semiring K] (m n : Nat) :
VectorArithmeticExecution K (m + n) :=
⟨fun _ => 0, 0, 0⟩A deterministic list containing every coefficient-index pair exactly once, in lexicographic row order.
def coefficientPairs (m n : Nat) : List (Fin m × Fin n) :=
(List.finRange m).flatMap fun i =>
(List.finRange n).map fun j => (i, j)The deterministic pair list enumerates the full Cartesian product.
theorem coefficientPairs_toFinset (m n : Nat) :
(coefficientPairs m n).toFinset =
(Finset.univ : Finset (Fin m)).product Finset.univ := by
ext ij
simp [coefficientPairs]Summing along the deterministic pair list is the usual nested finite sum over both index types.
private theorem coefficientPairs_sum [AddCommMonoid M] {m n : Nat}
(f : Fin m → Fin n → M) :
((coefficientPairs m n).map fun ij => f ij.1 ij.2).sum =
∑ i : Fin m, ∑ j : Fin n, f i j := by
have hflat (xs : List (Fin m)) :
(xs.flatMap fun i => (List.finRange n).map fun j => f i j).sum =
(xs.map fun i => ((List.finRange n).map fun j => f i j).sum).sum := by
induction xs with
| nil => simp
| cons i xs ih => simp [ih]
have hinner (i : Fin m) :
((List.finRange n).map fun j => f i j).sum =
∑ j : Fin n, f i j := by
rw [← List.sum_toFinset _ (List.nodup_finRange n)]
simp
rw [coefficientPairs]
simp only [List.map_flatMap, List.map_map, Function.comp_def]
rw [hflat]
simp_rw [hinner]
rw [← List.sum_toFinset _ (List.nodup_finRange m)]
simpCanonical schoolbook execution over every coefficient pair exactly once.
def schoolbookMulExec [Semiring K] {m n : Nat}
(a : CoeffVector K m) (b : CoeffVector K n) :
VectorArithmeticExecution K (m + n) :=
schoolbookPairsExec a b (coefficientPairs m n) (schoolbookInitial m n)The value returned by canonical schoolbook multiplication.
def schoolbookMul [Semiring K] {m n : Nat}
(a : CoeffVector K m) (b : CoeffVector K n) : CoeffVector K (m + n) :=
(schoolbookMulExec a b).valueFolding coefficient pairs adds their monomials to the initial polynomial.
private theorem schoolbookPairsExec_polynomial [CommSemiring K] {m n : Nat}
(a : CoeffVector K m) (b : CoeffVector K n)
(pairs : List (Fin m × Fin n))
(initial : VectorArithmeticExecution K (m + n)) :
vectorToPolynomial (schoolbookPairsExec a b pairs initial).value =
vectorToPolynomial initial.value +
(pairs.map fun ij => Polynomial.monomial (ij.1.1 + ij.2.1)
(a ij.1 * b ij.2)).sum := by
induction pairs generalizing initial with
| nil => simp [schoolbookPairsExec]
| cons ij pairs ih =>
simp only [schoolbookPairsExec, List.foldl]
change vectorToPolynomial
(schoolbookPairsExec a b pairs (schoolbookStep a b initial ij)).value =
vectorToPolynomial initial.value +
(List.map (fun ij => Polynomial.monomial (ij.1.1 + ij.2.1)
(a ij.1 * b ij.2)) (ij :: pairs)).sum
rw [ih]
simp [schoolbookStep, productIndex, vectorToPolynomial_addToBucket,
add_assoc]The full pair sum of coefficient monomials is the product of the two reconstructed input polynomials.
private theorem pairMonomialSum_eq [CommSemiring K] {m n : Nat}
(a : CoeffVector K m) (b : CoeffVector K n) :
(∑ i : Fin m, ∑ j : Fin n,
Polynomial.monomial (i.1 + j.1) (a i * b j)) =
vectorToPolynomial a * vectorToPolynomial b := by
rw [vectorToPolynomial, vectorToPolynomial]
simp only [Finset.sum_mul, Finset.mul_sum]
rw [Finset.sum_comm]
apply Finset.sum_congr rfl
intro i _
apply Finset.sum_congr rfl
intro j _
rw [Polynomial.monomial_mul_monomial]Schoolbook multiplication represents the product of the input polynomials.
theorem schoolbookMul_correct [CommSemiring K] {m n : Nat}
(a : CoeffVector K m) (b : CoeffVector K n) :
vectorToPolynomial (schoolbookMul a b) =
vectorToPolynomial a * vectorToPolynomial b := by
rw [schoolbookMul, schoolbookMulExec]
rw [schoolbookPairsExec_polynomial]
rw [show vectorToPolynomial (schoolbookInitial m n).value = 0 by
simp [schoolbookInitial, vectorToPolynomial]]
simp only [zero_add]
calc
((coefficientPairs m n).map fun ij =>
Polynomial.monomial (ij.1.1 + ij.2.1) (a ij.1 * b ij.2)).sum =
∑ i : Fin m, ∑ j : Fin n,
Polynomial.monomial (i.1 + j.1) (a i * b j) := by
simpa using (coefficientPairs_sum (m := m) (n := n)
(fun i j => Polynomial.monomial (i.1 + j.1) (a i * b j)))
_ = vectorToPolynomial a * vectorToPolynomial b :=
pairMonomialSum_eq a bA schoolbook output always fits its declared sum capacity.
theorem schoolbookMul_degreeBound [Semiring K] {m n : Nat}
(a : CoeffVector K m) (b : CoeffVector K n) :
(vectorToPolynomial (schoolbookMul a b)).degree < m + n :=
vectorToPolynomial_degree_lt _Pair folding increments the addition counter once per processed pair.
private theorem schoolbookPairsExec_additions [Semiring K] {m n : Nat}
(a : CoeffVector K m) (b : CoeffVector K n)
(pairs : List (Fin m × Fin n))
(initial : VectorArithmeticExecution K (m + n)) :
(schoolbookPairsExec a b pairs initial).additions =
initial.additions + pairs.length := by
induction pairs generalizing initial with
| nil => simp [schoolbookPairsExec]
| cons ij pairs ih =>
simp only [schoolbookPairsExec, List.foldl]
change (schoolbookPairsExec a b pairs
(schoolbookStep a b initial ij)).additions =
initial.additions + (ij :: pairs).length
rw [ih]
simp [schoolbookStep]
omegaPair folding increments the multiplication counter once per pair.
private theorem schoolbookPairsExec_multiplications [Semiring K] {m n : Nat}
(a : CoeffVector K m) (b : CoeffVector K n)
(pairs : List (Fin m × Fin n))
(initial : VectorArithmeticExecution K (m + n)) :
(schoolbookPairsExec a b pairs initial).multiplications =
initial.multiplications + pairs.length := by
induction pairs generalizing initial with
| nil => simp [schoolbookPairsExec]
| cons ij pairs ih =>
simp only [schoolbookPairsExec, List.foldl]
change (schoolbookPairsExec a b pairs
(schoolbookStep a b initial ij)).multiplications =
initial.multiplications + (ij :: pairs).length
rw [ih]
simp [schoolbookStep]
omegaSchoolbook execution performs exactly one addition per coefficient pair.
theorem schoolbookMulExec_additions [Semiring K] {m n : Nat}
(a : CoeffVector K m) (b : CoeffVector K n) :
(schoolbookMulExec a b).additions = m * n := by
rw [schoolbookMulExec, schoolbookPairsExec_additions]
simp [schoolbookInitial, coefficientPairs]Schoolbook execution performs exactly one multiplication per coefficient pair.
theorem schoolbookMulExec_multiplications [Semiring K] {m n : Nat}
(a : CoeffVector K m) (b : CoeffVector K n) :
(schoolbookMulExec a b).multiplications = m * n := by
rw [schoolbookMulExec, schoolbookPairsExec_multiplications]
simp [schoolbookInitial, coefficientPairs]Schoolbook execution performs exactly twice the number of coefficient pairs in charged arithmetic work.
theorem schoolbookMulWork_exact [Semiring K] {m n : Nat}
(a : CoeffVector K m) (b : CoeffVector K n) :
(schoolbookMulExec a b).work = 2 * (m * n) := by
rw [VectorArithmeticExecution.work, schoolbookMulExec_additions,
schoolbookMulExec_multiplications]
omegaend Chapter30end CLRSImports
import CLRSLean.FourthEdition.Chapter_30.Section_30_2_DFT_And_FFT.S1_RootsOfUnity
import CLRSLean.FourthEdition.Chapter_30.Section_30_2_DFT_And_FFT.S2_DFT
import CLRSLean.FourthEdition.Chapter_30.Section_30_2_DFT_And_FFT.S3_InversionAndConvolution
import CLRSLean.FourthEdition.Chapter_30.Section_30_2_DFT_And_FFT.RecursiveFFT
import CLRSLean.FourthEdition.Chapter_30.Section_30_2_DFT_And_FFT.RecursiveFFT.Definitions
import CLRSLean.FourthEdition.Chapter_30.Section_30_2_DFT_And_FFT.RecursiveFFT.Correctness
import CLRSLean.FourthEdition.Chapter_30.Section_30_2_DFT_And_FFT.RecursiveFFT.Costs
import CLRSLean.FourthEdition.Chapter_30.Section_30_2_DFT_And_FFT.PolynomialMultiplication30.2. The DFT and FFT
The generic algebraic core works over characteristic-zero fields: primitive
roots supply orthogonality, the positive-exponent CLRS transform is polynomial
evaluation at powers of a root, and idft gives both inverse directions.
Fourier-space pointwise multiplication is connected to cyclic convolution and,
under an explicit capacity premise, to ordinary polynomial multiplication.
The actual radix-2 FFT recursively splits even and odd coefficients, reuses a
successively generated twiddle trace, combines the children through a butterfly
layer, and supplies an inverse agreeing with idft. Its arithmetic counters
are projections of the same execution: each counter is exactly k * 2^k, and
zero padding lifts the total work to an all-input Theta(n log n) theorem.
complexDft_mathlib is the separate compatibility boundary with Mathlib's
complex ZMod.dft; its statement exposes the required output-index sign change.
The generic FFT multiplication execution is correct under the minimal no-wrap
capacity premise; the complex wrapper constructs a sufficient power-of-two
capacity and primitive root internally. The multiplication work field charges
three recursive transforms, pointwise products, and inverse scaling, with an
exact composition and an all-input Theta(n log n) bound. Section 30.3 builds
the iterative and layered-circuit refinements on this recursive core.
Implementation pages:
namespace CLRSnamespace Chapter30end Chapter30end CLRSDefinitions and proofs
CLRSLean.FourthEdition.Chapter_30.Section_30_2_DFT_And_FFT.S1_RootsOfUnity
Chapter 30.2: Roots-of-unity toolkit
This module packages the primitive-root facts used by the generic DFT, including the complete finite root-sum orthogonality statement.
namespace CLRSnamespace Chapter30
The first n powers of a primitive n-th root are injective.
theorem primitiveRoot_powers_injective [CommRing K] {n : Nat} {omega : K}
(homega : IsPrimitiveRoot omega n) :
Function.Injective (fun i : Fin n => omega ^ i.1) := by
intro i j hij
apply Fin.ext
exact homega.pow_inj i.2 j.2 hij
Squaring a primitive root of order 2 * n gives one of order n.
theorem primitiveRoot_square [CommMonoid K] {n : Nat} {omega : K}
(hn : 0 < n) (homega : IsPrimitiveRoot omega (2 * n)) :
IsPrimitiveRoot (omega ^ 2) n := by
exact homega.pow (by omega) rflThe inverse of a primitive root has the same order.
theorem primitiveRoot_inv [CommGroupWithZero K] {n : Nat} {omega : K}
(homega : IsPrimitiveRoot omega n) :
IsPrimitiveRoot omega⁻¹ n :=
homega.invThe halfway power of a primitive even-order root is negative one.
theorem primitiveRoot_half_pow_eq_neg_one [Field K] [CharZero K]
{n : Nat} (hn : 0 < n) {omega : K}
(homega : IsPrimitiveRoot omega (2 * n)) :
omega ^ n = -1 := by
have htwo : IsPrimitiveRoot (omega ^ n) 2 := by
exact homega.pow (by omega) (by omega)
exact htwo.eq_neg_one_of_two_right
A finite geometric sum vanishes when its base is a nontrivial n-th
root of one.
private theorem fin_root_sum_eq_zero [Field K] {n : Nat} {x : K}
(hxpow : x ^ n = 1) (hx : x ≠ 1) :
(∑ j : Fin n, x ^ j.1) = 0 := by
rw [Fin.sum_univ_eq_sum_range]
have hgeom := geom_sum_mul x n
have hzero : (∑ j ∈ Finset.range n, x ^ j) * (x - 1) = 0 := by
simpa [hxpow] using hgeom
exact (mul_eq_zero.mp hzero).resolve_right (sub_ne_zero.mpr hx)The powers of a primitive root sum to the cardinality precisely on exponents divisible by its order, and otherwise sum to zero.
theorem root_sum_orthogonality [Field K] [CharZero K]
{n exponent : Nat} (hn : 0 < n) {omega : K}
(homega : IsPrimitiveRoot omega n) :
(∑ j : Fin n, omega ^ (j.1 * exponent)) =
if n ∣ exponent then (n : K) else 0 := by
by_cases hdiv : n ∣ exponent
· rw [if_pos hdiv]
obtain ⟨t, rfl⟩ := hdiv
simp [pow_mul, homega.pow_eq_one, Nat.mul_assoc, Nat.mul_comm,
Nat.mul_left_comm]
· rw [if_neg hdiv]
have hxne : omega ^ exponent ≠ 1 := by
intro hpow
exact hdiv ((homega.pow_eq_one_iff_dvd exponent).mp hpow)
have hxpow : (omega ^ exponent) ^ n = 1 := by
rw [← pow_mul]
rw [Nat.mul_comm]
rw [pow_mul, homega.pow_eq_one, one_pow]
simpa [pow_mul, Nat.mul_comm] using
(fin_root_sum_eq_zero hxpow hxne)Orthogonality in the signed form used by Fourier inversion.
theorem root_sum_difference_orthogonality [Field K] [CharZero K]
{n : Nat} (hn : 0 < n) {omega : K}
(homega : IsPrimitiveRoot omega n) (i k : Fin n) :
(∑ j : Fin n,
omega ^ (j.1 * i.1) * omega⁻¹ ^ (j.1 * k.1)) =
if i = k then (n : K) else 0 := by
let x : K := omega ^ i.1 * omega⁻¹ ^ k.1
have homega0 : omega ≠ 0 := homega.ne_zero (Nat.ne_of_gt hn)
have hx_iff : x = 1 ↔ i = k := by
constructor
· intro hx
apply Fin.ext
apply homega.pow_inj i.2 k.2
have hk0 : omega ^ k.1 ≠ 0 := pow_ne_zero _ homega0
apply (div_eq_one_iff_eq hk0).mp
simpa [x, div_eq_mul_inv, inv_pow] using hx
· intro hik
subst k
simp [x, ← mul_pow, homega0]
have hxpow : x ^ n = 1 := by
have hi : (omega ^ i.1) ^ n = (omega ^ n) ^ i.1 := by
simp only [← pow_mul]
rw [Nat.mul_comm]
have hk : (omega⁻¹ ^ k.1) ^ n = (omega⁻¹ ^ n) ^ k.1 := by
simp only [← pow_mul]
rw [Nat.mul_comm]
dsimp [x]
rw [mul_pow, hi, hk, homega.pow_eq_one, homega.inv.pow_eq_one]
simp
have hsum :
(∑ j : Fin n,
omega ^ (j.1 * i.1) * omega⁻¹ ^ (j.1 * k.1)) =
∑ j : Fin n, x ^ j.1 := by
apply Finset.sum_congr rfl
intro j _
simp [x, mul_pow, pow_mul, Nat.mul_comm]
rw [hsum]
by_cases hik : i = k
· rw [if_pos hik, (hx_iff.mpr hik)]
simp
· rw [if_neg hik]
exact fin_root_sum_eq_zero hxpow (fun hx => hik (hx_iff.mp hx))end Chapter30end CLRSCLRSLean.FourthEdition.Chapter_30.Section_30_2_DFT_And_FFT.S2_DFT
Chapter 30.2: The generic discrete Fourier transform
The reusable transform is defined over an arbitrary semiring with the positive exponent convention used by CLRS. The complex compatibility theorem below keeps Mathlib's opposite sign visible at the boundary.
namespace CLRSnamespace Chapter30
The successive powers of omega, used as polynomial evaluation points.
def powerPoints [Monoid K] {n : Nat} (omega : K) : Fin n → K :=
fun k => omega ^ k.1The positive-exponent discrete Fourier transform used by CLRS.
def dft [Semiring K] {n : Nat} (omega : K) (a : CoeffVector K n) :
CoeffVector K n :=
fun k => ∑ j : Fin n, a j * omega ^ (j.1 * k.1)The DFT is precisely evaluation of the represented polynomial at powers of the root.
theorem dft_eq_pointValues [CommSemiring K] {n : Nat}
(omega : K) (a : CoeffVector K n) :
dft omega a = pointValues (powerPoints omega) (vectorToPolynomial a) := by
funext k
simp [dft, pointValues, powerPoints, vectorToPolynomial,
Polynomial.eval_finsetSum, pow_mul, Nat.mul_comm]The DFT maps the zero vector to zero.
theorem dft_zero [Semiring K] {n : Nat} (omega : K) :
dft omega (0 : CoeffVector K n) = 0 := by
funext k
simp [dft]The DFT preserves vector addition.
theorem dft_add [Semiring K] {n : Nat} (omega : K)
(a b : CoeffVector K n) :
dft omega (a + b) = dft omega a + dft omega b := by
funext k
simp [dft, add_mul, Finset.sum_add_distrib]Over a commutative semiring, the DFT commutes with scalar multiplication.
theorem dft_smul [CommSemiring K] {n : Nat} (omega c : K)
(a : CoeffVector K n) :
dft omega (c • a) = c • dft omega a := by
funext k
simp [dft, Finset.mul_sum, mul_assoc]
Transport a Fin n vector to Mathlib's ZMod n indexing.
def finVectorToZMod {n : Nat} [NeZero n] (a : CoeffVector ℂ n) :
ZMod n → ℂ :=
fun j => a ((ZMod.finEquiv n).symm j)
Transport a Mathlib ZMod n vector back to Fin n indexing.
def zmodVectorToFin {n : Nat} [NeZero n] (a : ZMod n → ℂ) :
CoeffVector ℂ n :=
fun j => a (ZMod.finEquiv n j)
Transporting a finite vector to ZMod and back is exact.
@[simp] theorem zmodVectorToFin_finVectorToZMod {n : Nat} [NeZero n]
(a : CoeffVector ℂ n) :
zmodVectorToFin (finVectorToZMod a) = a := by
funext j
simp [zmodVectorToFin, finVectorToZMod]
Transporting a ZMod vector to Fin and back is exact.
@[simp] theorem finVectorToZMod_zmodVectorToFin {n : Nat} [NeZero n]
(a : ZMod n → ℂ) :
finVectorToZMod (zmodVectorToFin a) = a := by
funext j
simp [zmodVectorToFin, finVectorToZMod]The ring equivalence sends a bounded natural index to its residue class.
theorem finEquiv_eq_natCast {n : Nat} [NeZero n] (j : Fin n) :
ZMod.finEquiv n j = (j.1 : ZMod n) := by
apply ZMod.val_injective
rw [ZMod.val_natCast, Nat.mod_eq_of_lt j.2]
cases n with
| zero => exact (NeZero.ne 0 rfl).elim
| succ n => rflA power of the positive principal root is Mathlib's standard additive character at the corresponding residue class.
private theorem principalRoot_pow_eq_stdAddChar {n : Nat} [NeZero n]
(m : Nat) :
(Complex.exp (2 * Real.pi * Complex.I / n)) ^ m =
ZMod.stdAddChar (m : ZMod n) := by
rw [← Complex.exp_nat_mul]
rw [show (m : ZMod n) = ((m : Int) : ZMod n) by norm_num]
rw [ZMod.stdAddChar_coe]
push_cast
congr 1
ringThe generic positive-sign transform agrees with Mathlib's negative-sign transform after explicitly negating the output index.
theorem complexDft_mathlib {n : Nat} [NeZero n]
(a : CoeffVector ℂ n) (k : Fin n) :
dft (Complex.exp (2 * Real.pi * Complex.I / n)) a k =
ZMod.dft (finVectorToZMod a) (-(ZMod.finEquiv n k)) := by
rw [ZMod.dft_apply]
rw [← (ZMod.finEquiv n).sum_comp]
apply Finset.sum_congr rfl
intro j _
have hj : (ZMod.finEquiv n).toEquiv j = (j.1 : ZMod n) :=
finEquiv_eq_natCast j
have hk : ZMod.finEquiv n k = (k.1 : ZMod n) :=
finEquiv_eq_natCast k
have hjback : (ZMod.finEquiv n).symm (j.1 : ZMod n) = j := by
rw [← hj]
exact Equiv.symm_apply_apply _ j
rw [hj, hk]
simp only [finVectorToZMod, neg_neg, mul_neg, neg_neg, smul_eq_mul]
rw [principalRoot_pow_eq_stdAddChar]
rw [hjback]
simpa only [Nat.cast_mul] using
(mul_comm (a j)
(ZMod.stdAddChar ((j.1 * k.1 : Nat) : ZMod n)))end Chapter30end CLRSCLRSLean.FourthEdition.Chapter_30.Section_30_2_DFT_And_FFT.S3_InversionAndConvolution
Chapter 30.2: Fourier inversion and convolution
This module proves the algebraic inverse of the generic DFT and then connects pointwise multiplication in Fourier space to cyclic convolution.
namespace CLRSnamespace Chapter30open PolynomialThe inverse DFT uses the inverse root and normalizes by the transform length.
def idft [Field K] {n : Nat} (omega : K) (a : CoeffVector K n) :
CoeffVector K n :=
fun k => (n : K)⁻¹ * dft omega⁻¹ a kA positive natural has nonzero image in a characteristic-zero field.
theorem natCast_ne_zero_of_pos [Field K] [CharZero K]
{n : Nat} (hn : 0 < n) :
(n : K) ≠ 0 := by
exact_mod_cast Nat.ne_of_gt hnApplying the inverse transform after the forward transform recovers the input vector.
theorem idft_dft [Field K] [CharZero K] {n : Nat} (hn : 0 < n)
{omega : K} (homega : IsPrimitiveRoot omega n)
(a : CoeffVector K n) :
idft omega (dft omega a) = a := by
funext k
change (n : K)⁻¹ *
(∑ i : Fin n, (∑ j : Fin n,
a j * omega ^ (j.1 * i.1)) * omega⁻¹ ^ (i.1 * k.1)) = a k
calc
(n : K)⁻¹ *
(∑ i : Fin n, (∑ j : Fin n,
a j * omega ^ (j.1 * i.1)) * omega⁻¹ ^ (i.1 * k.1)) =
(n : K)⁻¹ *
(∑ i : Fin n, ∑ j : Fin n,
(a j * omega ^ (j.1 * i.1)) *
omega⁻¹ ^ (i.1 * k.1)) := by
congr 1
apply Finset.sum_congr rfl
intro i _
rw [Finset.sum_mul]
_ = (n : K)⁻¹ *
(∑ j : Fin n, ∑ i : Fin n,
(a j * omega ^ (j.1 * i.1)) *
omega⁻¹ ^ (i.1 * k.1)) := by
rw [Finset.sum_comm]
_ = (n : K)⁻¹ *
(∑ j : Fin n, a j *
(∑ i : Fin n,
omega ^ (i.1 * j.1) * omega⁻¹ ^ (i.1 * k.1))) := by
congr 1
apply Finset.sum_congr rfl
intro j _
rw [Finset.mul_sum]
apply Finset.sum_congr rfl
intro i _
rw [Nat.mul_comm j.1 i.1]
ring
_ = (n : K)⁻¹ *
(∑ j : Fin n, a j * (if j = k then (n : K) else 0)) := by
congr 1
apply Finset.sum_congr rfl
intro j _
rw [root_sum_difference_orthogonality hn homega j k]
_ = a k := by
simp only [mul_ite, mul_zero]
rw [Finset.sum_ite_eq' Finset.univ k,
if_pos (Finset.mem_univ k)]
calc
(n : K)⁻¹ * (a k * (n : K)) =
a k * ((n : K)⁻¹ * (n : K)) := by ring
_ = a k := by
rw [inv_mul_cancel₀ (natCast_ne_zero_of_pos hn), mul_one]Applying the forward transform after the inverse transform also recovers the input vector.
theorem dft_idft [Field K] [CharZero K] {n : Nat} (hn : 0 < n)
{omega : K} (homega : IsPrimitiveRoot omega n)
(a : CoeffVector K n) :
dft omega (idft omega a) = a := by
have hinv := idft_dft hn (primitiveRoot_inv homega) a
have hinv' :
(n : K)⁻¹ • dft omega (dft omega⁻¹ a) = a := by
calc
(n : K)⁻¹ • dft omega (dft omega⁻¹ a) =
idft omega⁻¹ (dft omega⁻¹ a) := by
funext k
simp only [idft, inv_inv, Pi.smul_apply, smul_eq_mul]
_ = a := hinv
have hidft : idft omega a = (n : K)⁻¹ • dft omega⁻¹ a := by
funext k
rfl
calc
dft omega (idft omega a) =
dft omega ((n : K)⁻¹ • dft omega⁻¹ a) := by rw [hidft]
_ = (n : K)⁻¹ • dft omega (dft omega⁻¹ a) :=
dft_smul omega (n : K)⁻¹ (dft omega⁻¹ a)
_ = a := hinv'A primitive-root DFT of positive length is injective.
theorem dft_injective [Field K] [CharZero K] {n : Nat} (hn : 0 < n)
{omega : K} (homega : IsPrimitiveRoot omega n) :
Function.Injective
(dft omega : CoeffVector K n → CoeffVector K n) := by
intro a b h
calc
a = idft omega (dft omega a) := (idft_dft hn homega a).symm
_ = idft omega (dft omega b) := congrArg (idft omega) h
_ = b := idft_dft hn homega bTotal subtraction modulo a positive vector length.
def cyclicSub {n : Nat} (hn : 0 < n) (k j : Fin n) : Fin n :=
⟨(k.1 + n - j.1) % n, Nat.mod_lt _ hn⟩Cyclic convolution of two fixed-capacity vectors.
def cyclicConvolution [Semiring K] {n : Nat} (hn : 0 < n)
(a b : CoeffVector K n) : CoeffVector K n :=
fun k => ∑ j : Fin n, a j * b (cyclicSub hn k j)
The concrete modular subtraction agrees with subtraction in ZMod n.
private theorem finEquiv_cyclicSub {n : Nat} [NeZero n] (hn : 0 < n)
(k j : Fin n) :
ZMod.finEquiv n (cyclicSub hn k j) =
ZMod.finEquiv n k - ZMod.finEquiv n j := by
rw [show cyclicSub hn k j =
(ZMod.finEquiv n).symm
(ZMod.finEquiv n k - ZMod.finEquiv n j) by
cases n with
| zero => omega
| succ n =>
apply Fin.ext
simp only [cyclicSub, ZMod.finEquiv]
congr 1
omega]
exact Equiv.apply_symm_apply _ _
For fixed j, subtracting j modulo n permutes every index.
private def cyclicSubEquiv {n : Nat} [NeZero n] (j : Fin n) :
Fin n ≃ Fin n :=
(ZMod.finEquiv n).toEquiv.trans
((Equiv.subRight (ZMod.finEquiv n j)).trans
(ZMod.finEquiv n).toEquiv.symm)The cyclic-subtraction equivalence computes the concrete modular index.
private theorem cyclicSubEquiv_apply {n : Nat} [NeZero n]
(hn : 0 < n) (k j : Fin n) :
cyclicSubEquiv j k = cyclicSub hn k j := by
apply (ZMod.finEquiv n).injective
rw [finEquiv_cyclicSub hn k j]
simp [cyclicSubEquiv]
A cyclicly subtracted index reconstructs the original index modulo n.
private theorem cyclicSub_add_modEq {n : Nat} [NeZero n]
(hn : 0 < n) (k j : Fin n) :
k.1 ≡ j.1 + (cyclicSub hn k j).1 [MOD n] := by
rw [← ZMod.natCast_eq_natCast_iff]
simp only [Nat.cast_add]
rw [← finEquiv_eq_natCast k, ← finEquiv_eq_natCast j,
← finEquiv_eq_natCast (cyclicSub hn k j),
finEquiv_cyclicSub hn k j]
abelA primitive root turns modular index reconstruction into multiplication of powers.
private theorem pow_index_eq_mul_pow_cyclicSub [Field K]
{n : Nat} [NeZero n] (hn : 0 < n) {omega : K}
(homega : IsPrimitiveRoot omega n) (k j r : Fin n) :
omega ^ (k.1 * r.1) =
omega ^ (j.1 * r.1) *
omega ^ ((cyclicSub hn k j).1 * r.1) := by
have hmod :
k.1 * r.1 ≡
j.1 * r.1 + (cyclicSub hn k j).1 * r.1 [MOD n] := by
simpa only [Nat.add_mul] using
(cyclicSub_add_modEq hn k j).mul_right r.1
rw [← pow_add]
exact pow_eq_pow_of_modEq hmod homega.pow_eq_oneReindexing by cyclic subtraction separates one Fourier kernel factor.
private theorem sum_cyclicSub_mul_pow [Field K] {n : Nat} (hn : 0 < n)
{omega : K} (homega : IsPrimitiveRoot omega n)
(b : CoeffVector K n) (j r : Fin n) :
(∑ k : Fin n,
b (cyclicSub hn k j) * omega ^ (k.1 * r.1)) =
omega ^ (j.1 * r.1) *
(∑ l : Fin n, b l * omega ^ (l.1 * r.1)) := by
letI : NeZero n := ⟨Nat.ne_of_gt hn⟩
let e : Fin n ≃ Fin n := cyclicSubEquiv j
calc
(∑ k : Fin n,
b (cyclicSub hn k j) * omega ^ (k.1 * r.1)) =
∑ k : Fin n, b (e k) * omega ^ (k.1 * r.1) := by
apply Finset.sum_congr rfl
intro k _
rw [cyclicSubEquiv_apply hn k j]
_ = ∑ l : Fin n,
b l * omega ^ ((e.symm l).1 * r.1) := by
simpa only [Equiv.symm_apply_apply] using
(e.sum_comp
(fun l : Fin n =>
b l * omega ^ ((e.symm l).1 * r.1)))
_ = ∑ l : Fin n,
b l * (omega ^ (j.1 * r.1) * omega ^ (l.1 * r.1)) := by
apply Finset.sum_congr rfl
intro l _
have he : cyclicSub hn (e.symm l) j = l := by
rw [← cyclicSubEquiv_apply hn (e.symm l) j]
exact e.apply_symm_apply l
rw [pow_index_eq_mul_pow_cyclicSub hn homega (e.symm l) j r, he]
_ = omega ^ (j.1 * r.1) *
(∑ l : Fin n, b l * omega ^ (l.1 * r.1)) := by
rw [Finset.mul_sum]
apply Finset.sum_congr rfl
intro l _
ringThe DFT sends cyclic convolution to pointwise multiplication.
theorem dft_cyclicConvolution [Field K] [CharZero K]
{n : Nat} (hn : 0 < n) {omega : K}
(homega : IsPrimitiveRoot omega n)
(a b : CoeffVector K n) :
dft omega (cyclicConvolution hn a b) =
pointwiseMul (dft omega a) (dft omega b) := by
funext r
change (∑ k : Fin n,
(∑ j : Fin n, a j * b (cyclicSub hn k j)) *
omega ^ (k.1 * r.1)) =
(∑ j : Fin n, a j * omega ^ (j.1 * r.1)) *
(∑ l : Fin n, b l * omega ^ (l.1 * r.1))
calc
(∑ k : Fin n,
(∑ j : Fin n, a j * b (cyclicSub hn k j)) *
omega ^ (k.1 * r.1)) =
∑ k : Fin n, ∑ j : Fin n,
(a j * b (cyclicSub hn k j)) *
omega ^ (k.1 * r.1) := by
apply Finset.sum_congr rfl
intro k _
rw [Finset.sum_mul]
_ = ∑ j : Fin n, ∑ k : Fin n,
(a j * b (cyclicSub hn k j)) *
omega ^ (k.1 * r.1) := by
rw [Finset.sum_comm]
_ = ∑ j : Fin n, a j *
(∑ k : Fin n,
b (cyclicSub hn k j) * omega ^ (k.1 * r.1)) := by
apply Finset.sum_congr rfl
intro j _
rw [Finset.mul_sum]
apply Finset.sum_congr rfl
intro k _
ring
_ = ∑ j : Fin n, a j *
(omega ^ (j.1 * r.1) *
(∑ l : Fin n, b l * omega ^ (l.1 * r.1))) := by
apply Finset.sum_congr rfl
intro j _
rw [sum_cyclicSub_mul_pow hn homega b j r]
_ = (∑ j : Fin n, a j * omega ^ (j.1 * r.1)) *
(∑ l : Fin n, b l * omega ^ (l.1 * r.1)) := by
rw [Finset.sum_mul]
apply Finset.sum_congr rfl
intro j _
ringInverse-transforming pointwise products recovers cyclic convolution.
theorem idft_pointwiseMul [Field K] [CharZero K]
{n : Nat} (hn : 0 < n) {omega : K}
(homega : IsPrimitiveRoot omega n)
(a b : CoeffVector K n) :
idft omega (pointwiseMul (dft omega a) (dft omega b)) =
cyclicConvolution hn a b := by
rw [← dft_cyclicConvolution hn homega, idft_dft hn homega]When a polynomial product fits in the vector capacity, cyclic convolution has no wrapped contribution and equals its coefficient vector.
theorem cyclicConvolution_eq_coeffVector_mul [Field K]
{n : Nat} (hn : 0 < n) (p q : K[X])
(hfit : (p * q).degree < n) :
cyclicConvolution hn (coeffVector n p) (coeffVector n q) =
coeffVector n (p * q) := by
by_cases hp : p = 0
· subst p
funext k
simp [cyclicConvolution, coeffVector]
by_cases hq : q = 0
· subst q
funext k
simp [cyclicConvolution, coeffVector]
have hpq : p * q ≠ 0 := mul_ne_zero hp hq
have hnatfit : p.natDegree + q.natDegree < n := by
rw [← Polynomial.natDegree_mul hp hq]
exact (Polynomial.natDegree_lt_iff_degree_lt hpq).mpr hfit
funext k
change (∑ j : Fin n,
p.coeff j.1 * q.coeff (cyclicSub hn k j).1) =
(p * q).coeff k.1
rw [show (∑ j : Fin n,
p.coeff j.1 * q.coeff (cyclicSub hn k j).1) =
∑ j ∈ Finset.range n,
p.coeff j * q.coeff
(cyclicSub hn k ⟨j % n, Nat.mod_lt _ hn⟩).1 by
simpa [Nat.mod_eq_of_lt] using
(Fin.sum_univ_eq_sum_range
(fun j : Nat =>
p.coeff j * q.coeff
(cyclicSub hn k ⟨j % n, Nat.mod_lt _ hn⟩).1) n)]
rw [Polynomial.coeff_mul,
Finset.Nat.sum_antidiagonal_eq_sum_range_succ
(fun i j => p.coeff i * q.coeff j) k.1]
calc
(∑ j ∈ Finset.range n,
p.coeff j *
q.coeff (cyclicSub hn k
⟨j % n, Nat.mod_lt _ hn⟩).1) =
∑ j ∈ Finset.range (k.1 + 1),
p.coeff j *
q.coeff (cyclicSub hn k
⟨j % n, Nat.mod_lt _ hn⟩).1 := by
symm
apply Finset.sum_subset
· intro j hj
simp only [Finset.mem_range] at hj ⊢
omega
· intro j hjn hjk
simp only [Finset.mem_range] at hjn hjk
have hkj : k.1 < j := by omega
have hjcast :
(⟨j % n, Nat.mod_lt _ hn⟩ : Fin n) = ⟨j, hjn⟩ := by
apply Fin.ext
exact Nat.mod_eq_of_lt hjn
have hcyc :
(cyclicSub hn k
⟨j % n, Nat.mod_lt _ hn⟩).1 = k.1 + n - j := by
rw [hjcast]
simp only [cyclicSub]
rw [Nat.mod_eq_of_lt (by omega)]
by_contra hterm
have hpcoeff : p.coeff j ≠ 0 := by
intro hz
exact hterm (by simp [hz])
have hqcoeff :
q.coeff (cyclicSub hn k
⟨j % n, Nat.mod_lt _ hn⟩).1 ≠ 0 := by
intro hz
exact hterm (by simp [hz])
have hjdeg : j ≤ p.natDegree :=
Polynomial.le_natDegree_of_ne_zero hpcoeff
have hcycdeg :
(cyclicSub hn k
⟨j % n, Nat.mod_lt _ hn⟩).1 ≤ q.natDegree :=
Polynomial.le_natDegree_of_ne_zero hqcoeff
omega
_ = ∑ j ∈ Finset.range (k.1 + 1),
p.coeff j * q.coeff (k.1 - j) := by
apply Finset.sum_congr rfl
intro j hj
simp only [Finset.mem_range] at hj
have hjn : j < n := by omega
have hjk : j ≤ k.1 := by omega
have hjcast :
(⟨j % n, Nat.mod_lt _ hn⟩ : Fin n) = ⟨j, hjn⟩ := by
apply Fin.ext
exact Nat.mod_eq_of_lt hjn
congr 2
rw [hjcast]
simp only [cyclicSub]
rw [show k.1 + n - j = n + (k.1 - j) by omega]
have hlt : k.1 - j < n :=
lt_of_le_of_lt (Nat.sub_le _ _) k.2
simp [Nat.add_mod, Nat.mod_eq_of_lt hlt]end Chapter30end CLRSCLRSLean.FourthEdition.Chapter_30.Section_30_2_DFT_And_FFT.RecursiveFFT.Definitions
Chapter 30.2: Executable recursive radix-2 FFT
The execution below generates twiddles successively, reuses that same trace to obtain the child root, and records arithmetic operations in the value-producing recursion.
namespace CLRSnamespace Chapter30A successor power-of-two length splits into two equal halves.
private theorem powTwo_succ_split (k : Nat) :
2 ^ (k + 1) = 2 ^ k + 2 ^ k := by
rw [pow_succ]
omegaEmbed a half-size index as an even full-size coefficient index.
def evenIndex {k : Nat} (i : Fin (2 ^ k)) : Fin (2 ^ (k + 1)) :=
⟨2 * i.1, by
have := i.2
simp [pow_succ]
omega⟩Embed a half-size index as an odd full-size coefficient index.
def oddIndex {k : Nat} (i : Fin (2 ^ k)) : Fin (2 ^ (k + 1)) :=
⟨2 * i.1 + 1, by
have := i.2
simp [pow_succ]
omega⟩The natural value of an embedded even coefficient index.
@[simp] theorem evenIndex_val {k : Nat} (i : Fin (2 ^ k)) :
(evenIndex i).1 = 2 * i.1 := rflThe natural value of an embedded odd coefficient index.
@[simp] theorem oddIndex_val {k : Nat} (i : Fin (2 ^ k)) :
(oddIndex i).1 = 2 * i.1 + 1 := rflThe even-indexed coefficient half.
def evenCoeffs {K : Type*} {k : Nat} (a : PowTwoVec K (k + 1)) :
PowTwoVec K k := fun i => a (evenIndex i)The odd-indexed coefficient half.
def oddCoeffs {K : Type*} {k : Nat} (a : PowTwoVec K (k + 1)) :
PowTwoVec K k := fun i => a (oddIndex i)Reading the even coefficient half uses the corresponding even index.
@[simp] theorem evenCoeffs_apply {K : Type*} {k : Nat}
(a : PowTwoVec K (k + 1)) (i : Fin (2 ^ k)) :
evenCoeffs a i = a (evenIndex i) := rflReading the odd coefficient half uses the corresponding odd index.
@[simp] theorem oddCoeffs_apply {K : Type*} {k : Nat}
(a : PowTwoVec K (k + 1)) (i : Fin (2 ^ k)) :
oddCoeffs a i = a (oddIndex i) := rflReassociate the successor power-of-two index type into two halves.
def powTwoSuccEquiv (k : Nat) :
Fin (2 ^ (k + 1)) ≃ Fin (2 ^ k + 2 ^ k) :=
finCongr (powTwo_succ_split k)Index into the lower half of a successor power-of-two vector.
def lowerHalfIndex {k : Nat} (i : Fin (2 ^ k)) : Fin (2 ^ (k + 1)) :=
(powTwoSuccEquiv k).symm (Fin.castAdd (2 ^ k) i)Index into the upper half of a successor power-of-two vector.
def upperHalfIndex {k : Nat} (i : Fin (2 ^ k)) : Fin (2 ^ (k + 1)) :=
(powTwoSuccEquiv k).symm (Fin.natAdd (2 ^ k) i)The lower-half embedding preserves the natural index value.
@[simp] theorem lowerHalfIndex_val {k : Nat} (i : Fin (2 ^ k)) :
(lowerHalfIndex i).1 = i.1 := by
simp [lowerHalfIndex, powTwoSuccEquiv]The upper-half embedding offsets the natural index by the half length.
@[simp] theorem upperHalfIndex_val {k : Nat} (i : Fin (2 ^ k)) :
(upperHalfIndex i).1 = 2 ^ k + i.1 := by
simp [upperHalfIndex, powTwoSuccEquiv]
omegaJoin two equal power-of-two halves without leaking casts downstream.
def joinHalves {K : Type*} {k : Nat} (lower upper : PowTwoVec K k) :
PowTwoVec K (k + 1) :=
fun i => Fin.append lower upper (powTwoSuccEquiv k i)Reading the lower embedding of joined vectors returns the lower input.
@[simp] theorem joinHalves_lower {K : Type*} {k : Nat}
(lower upper : PowTwoVec K k) (i : Fin (2 ^ k)) :
joinHalves lower upper (lowerHalfIndex i) = lower i := by
simp [joinHalves, lowerHalfIndex, powTwoSuccEquiv]Reading the upper embedding of joined vectors returns the upper input.
@[simp] theorem joinHalves_upper {K : Type*} {k : Nat}
(lower upper : PowTwoVec K k) (i : Fin (2 ^ k)) :
joinHalves lower upper (upperHalfIndex i) = upper i := by
change Fin.append lower upper (Fin.natAdd (2 ^ k) i) = upper i
exact Fin.append_right lower upper i
A successive twiddle-generation execution. next is the accumulator
after the last charged multiplication.
structure TwiddleExecution (K : Type*) where
value : List K
next : K
multiplications : Nat
Generate n successive values, charging every accumulator update.
def twiddlePowersAuxExec [Monoid K] (omega : K) :
Nat → K → TwiddleExecution K
| 0, current => ⟨[], current, 0⟩
| n + 1, current =>
let child := twiddlePowersAuxExec omega n (current * omega)
⟨current :: child.value, child.next, child.multiplications + 1⟩Value projection of the successive twiddle generator.
def twiddlePowersAux [Monoid K] (omega : K) (n : Nat) (current : K) : List K :=
(twiddlePowersAuxExec omega n current).valueSuccessive twiddle generation returns exactly the requested number of values.
theorem twiddlePowersAuxExec_length [Monoid K]
(omega : K) (n : Nat) (current : K) :
(twiddlePowersAuxExec omega n current).value.length = n := by
induction n generalizing current with
| zero => rfl
| succ n ih => simp [twiddlePowersAuxExec, ih]The value-only twiddle generator has the requested length.
theorem twiddlePowersAux_length [Monoid K]
(omega : K) (n : Nat) (current : K) :
(twiddlePowersAux omega n current).length = n :=
twiddlePowersAuxExec_length omega n currentSuccessive twiddle generation charges one multiplication per output.
theorem twiddlePowersAuxExec_multiplications [Monoid K]
(omega : K) (n : Nat) (current : K) :
(twiddlePowersAuxExec omega n current).multiplications = n := by
induction n generalizing current with
| zero => rfl
| succ n ih => simp [twiddlePowersAuxExec, ih]
The final twiddle accumulator is the initial value times omega ^ n.
private theorem twiddlePowersAuxExec_next [Monoid K]
(omega : K) (n : Nat) (current : K) :
(twiddlePowersAuxExec omega n current).next = current * omega ^ n := by
induction n generalizing current with
| zero => simp [twiddlePowersAuxExec]
| succ n ih =>
simp [twiddlePowersAuxExec, ih, pow_succ', mul_assoc]Every generated twiddle equals the corresponding successive power.
private theorem twiddlePowersAuxExec_get [Monoid K]
(omega : K) (n : Nat) (current : K) (i : Nat) (hi : i < n) :
(twiddlePowersAuxExec omega n current).value.getD i current =
current * omega ^ i := by
induction n generalizing current i with
| zero => omega
| succ n ih =>
cases i with
| zero => simp [twiddlePowersAuxExec]
| succ i =>
simp only [twiddlePowersAuxExec, List.getD_cons_succ]
have hlen :
(twiddlePowersAuxExec omega n (current * omega)).value.length = n :=
twiddlePowersAuxExec_length omega n (current * omega)
have hi' :
i < (twiddlePowersAuxExec omega n (current * omega)).value.length := by
omega
rw [List.getD_eq_getElem _ _ hi']
rw [← List.getD_eq_getElem _ (current * omega) hi']
rw [ih (current * omega) i (by omega)]
simp [pow_succ', mul_assoc]Convert a checked twiddle trace to a fixed-capacity vector.
def twiddleVectorOfExecution {K : Type*} (run : TwiddleExecution K) {n : Nat}
(hlen : run.value.length = n) : CoeffVector K n :=
fun i => run.value.get (Fin.cast hlen.symm i)
The first n powers, obtained from one successive generator execution.
def twiddlePowers [Monoid K] (omega : K) (n : Nat) : CoeffVector K n :=
let run := twiddlePowersAuxExec omega n 1
twiddleVectorOfExecution run (twiddlePowersAuxExec_length omega n 1)
The checked twiddle vector contains the first successive powers of omega.
theorem twiddlePowers_eq_pow [Monoid K]
(omega : K) (n : Nat) (i : Fin n) :
twiddlePowers omega n i = omega ^ i.1 := by
change (twiddlePowersAuxExec omega n (1 : K)).value.get
(Fin.cast (twiddlePowersAuxExec_length omega n 1).symm i) = omega ^ i.1
rw [← List.getD_eq_get
(twiddlePowersAuxExec omega n (1 : K)).value 1
(Fin.cast (twiddlePowersAuxExec_length omega n 1).symm i)]
simpa using twiddlePowersAuxExec_get omega n (1 : K) i.1 i.2Recover the squared child root from the charged twiddle trace.
def twiddleChildRoot [One K] (k : Nat) (_omega : K)
(run : TwiddleExecution K) : K :=
if k = 0 then 1 else run.value.getD 2 run.nextA positive-size twiddle trace exposes the squared root used by child FFTs.
theorem twiddleChildRoot_eq_square [Monoid K] {k : Nat} (hk : 0 < k)
(omega : K) :
twiddleChildRoot k omega (twiddlePowersAuxExec omega (2 ^ k) 1) =
omega ^ 2 := by
rw [twiddleChildRoot, if_neg (Nat.ne_of_gt hk)]
by_cases hlen : 2 < 2 ^ k
· have hactual :
2 < (twiddlePowersAuxExec omega (2 ^ k) (1 : K)).value.length := by
simpa [twiddlePowersAuxExec_length] using hlen
rw [List.getD_eq_getElem _ _ hactual]
rw [← List.getD_eq_getElem _ (1 : K) hactual]
rw [twiddlePowersAuxExec_get omega (2 ^ k) (1 : K) 2 hlen]
simp
· have hk_le_one : k ≤ 1 := by
by_contra hnot
have htwo : 2 ≤ k := by omega
have hp : 2 ^ 2 ≤ 2 ^ k :=
Nat.pow_le_pow_right (by omega) htwo
norm_num at hp
omega
have hkone : k = 1 := by omega
subst k
simp [twiddlePowersAuxExec, pow_two]One radix-2 butterfly layer and its shared-arithmetic charge counters.
structure ButterflyExecution (K : Type*) (k : Nat) where
value : PowTwoVec K (k + 1)
addSubtractions : Nat
multiplications : Nat
Consume a previously evaluated twiddle trace in a butterfly layer.
The charge counts w j * v j once per butterfly, shared by the sum and
difference. The function-valued output below does not itself memoize this
product across separate evaluations of its two output slots.
def butterflyLayerFromTwiddleExec [Ring K] {k : Nat} (omega : K)
(twiddleRun : TwiddleExecution K)
(hrun : twiddleRun = twiddlePowersAuxExec omega (2 ^ k) 1)
(u v : PowTwoVec K k) : ButterflyExecution K k :=
let w := twiddleVectorOfExecution twiddleRun
(by simpa [hrun] using twiddlePowersAuxExec_length omega (2 ^ k) 1)
⟨joinHalves
(fun j => u j + w j * v j)
(fun j => u j - w j * v j),
2 * 2 ^ k,
2 ^ k + twiddleRun.multiplications⟩Execute one standalone butterfly layer, including twiddle generation.
def butterflyLayerExec [Ring K] {k : Nat} (omega : K)
(u v : PowTwoVec K k) : ButterflyExecution K k :=
let twiddleRun := twiddlePowersAuxExec omega (2 ^ k) 1
butterflyLayerFromTwiddleExec omega twiddleRun rfl u vValue projection of one butterfly execution.
def butterflyLayer [Ring K] {k : Nat} (omega : K)
(u v : PowTwoVec K k) : PowTwoVec K (k + 1) :=
(butterflyLayerExec omega u v).valueThe canonical checked execution vector is the public twiddle vector.
private theorem twiddleVectorOfExecution_canonical [Monoid K]
(omega : K) (n : Nat) :
twiddleVectorOfExecution (twiddlePowersAuxExec omega n 1)
(twiddlePowersAuxExec_length omega n 1) = twiddlePowers omega n := rflThe lower butterfly output is the sum with its twiddle product.
@[simp] theorem butterflyLayer_lower [Ring K] {k : Nat} (omega : K)
(u v : PowTwoVec K k) (j : Fin (2 ^ k)) :
butterflyLayer omega u v (lowerHalfIndex j) =
u j + omega ^ j.1 * v j := by
simp [butterflyLayer, butterflyLayerExec, butterflyLayerFromTwiddleExec,
twiddleVectorOfExecution_canonical, twiddlePowers_eq_pow]The upper butterfly output is the difference with its twiddle product.
@[simp] theorem butterflyLayer_upper [Ring K] {k : Nat} (omega : K)
(u v : PowTwoVec K k) (j : Fin (2 ^ k)) :
butterflyLayer omega u v (upperHalfIndex j) =
u j - omega ^ j.1 * v j := by
simp [butterflyLayer, butterflyLayerExec, butterflyLayerFromTwiddleExec,
twiddleVectorOfExecution_canonical, twiddlePowers_eq_pow]A butterfly layer charges two additions/subtractions per local offset.
theorem butterflyLayerExec_addSubtractions [Ring K] {k : Nat} (omega : K)
(u v : PowTwoVec K k) :
(butterflyLayerExec omega u v).addSubtractions = 2 * 2 ^ k := rflA butterfly layer charges data products and successive twiddle updates.
theorem butterflyLayerExec_multiplications [Ring K] {k : Nat} (omega : K)
(u v : PowTwoVec K k) :
(butterflyLayerExec omega u v).multiplications = 2 * 2 ^ k := by
simp [butterflyLayerExec, butterflyLayerFromTwiddleExec,
twiddlePowersAuxExec_multiplications]
omegaResult and counters of the canonical recursive FFT.
structure FFTExecution (K : Type*) (k : Nat) where
value : PowTwoVec K k
addSubtractions : Nat
multiplications : NatTotal charged arithmetic operations.
def FFTExecution.work (r : FFTExecution K k) : Nat :=
r.addSubtractions + r.multiplicationsThe canonical radix-2 execution. Its child root is extracted from the same twiddle trace consumed by the butterfly.
def recursiveFFTExec [Ring K] :
{k : Nat} → K → PowTwoVec K k → FFTExecution K k
| 0, _, a => ⟨a, 0, 0⟩
| k + 1, omega, a =>
let twiddleRun := twiddlePowersAuxExec omega (2 ^ k) 1
let childRoot := twiddleChildRoot k omega twiddleRun
let evenRun := recursiveFFTExec childRoot (evenCoeffs a)
let oddRun := recursiveFFTExec childRoot (oddCoeffs a)
let layer := butterflyLayerFromTwiddleExec omega twiddleRun rfl
evenRun.value oddRun.value
⟨layer.value,
evenRun.addSubtractions + oddRun.addSubtractions + layer.addSubtractions,
evenRun.multiplications + oddRun.multiplications + layer.multiplications⟩Value projection of the canonical recursive execution.
def recursiveFFT [Ring K] {k : Nat} (omega : K) (a : PowTwoVec K k) :
PowTwoVec K k :=
(recursiveFFTExec omega a).valueErasing the counters of a recursive execution yields the public FFT value.
theorem recursiveFFTExec_value [Ring K] {k : Nat}
(omega : K) (a : PowTwoVec K k) :
(recursiveFFTExec omega a).value = recursiveFFT omega a := rflRecursive inverse FFT: inverse-root recursive FFT followed by scaling.
def recursiveIFFT [Field K] {k : Nat} (omega : K) (a : PowTwoVec K k) :
PowTwoVec K k :=
fun i => ((2 ^ k : Nat) : K)⁻¹ * recursiveFFT omega⁻¹ a iend Chapter30end CLRSCLRSLean.FourthEdition.Chapter_30.Section_30_2_DFT_And_FFT.RecursiveFFT.Correctness
Chapter 30.2: Recursive FFT correctness
namespace CLRSnamespace Chapter30open PolynomialInterleaved even/odd indices form an equivalence with a successor power-of-two index.
private def radixTwoIndexEquiv (k : Nat) :
Fin (2 ^ k) × Fin 2 ≃ Fin (2 ^ (k + 1)) :=
finProdFinEquiv.trans (finCongr (by simp [pow_succ, Nat.mul_comm]))The zero branch of radix-two reindexing has the expected even value.
private theorem radixTwoIndexEquiv_zero {k : Nat} (j : Fin (2 ^ k)) :
(radixTwoIndexEquiv k (j, (0 : Fin 2))).1 = 2 * j.1 := by
simp [radixTwoIndexEquiv, finProdFinEquiv, Nat.mul_comm]The one branch of radix-two reindexing has the expected odd value.
private theorem radixTwoIndexEquiv_one {k : Nat} (j : Fin (2 ^ k)) :
(radixTwoIndexEquiv k (j, (1 : Fin 2))).1 = 2 * j.1 + 1 := by
simp [radixTwoIndexEquiv, finProdFinEquiv, Nat.mul_comm,
Nat.add_comm]
The zero branch of radix-two reindexing is evenIndex.
private theorem radixTwoIndexEquiv_zero_eq {k : Nat} (j : Fin (2 ^ k)) :
radixTwoIndexEquiv k (j, (0 : Fin 2)) = evenIndex j := by
apply Fin.ext
exact radixTwoIndexEquiv_zero j
The one branch of radix-two reindexing is oddIndex.
private theorem radixTwoIndexEquiv_one_eq {k : Nat} (j : Fin (2 ^ k)) :
radixTwoIndexEquiv k (j, (1 : Fin 2)) = oddIndex j := by
apply Fin.ext
simpa [Nat.add_comm] using radixTwoIndexEquiv_one jThe coefficient polynomial splits into its even and odd parts.
theorem polynomial_evenOdd_split [CommRing K] {k : Nat}
(a : PowTwoVec K (k + 1)) :
vectorToPolynomial a =
(vectorToPolynomial (evenCoeffs a)).comp (Polynomial.X ^ 2) +
Polynomial.X *
(vectorToPolynomial (oddCoeffs a)).comp (Polynomial.X ^ 2) := by
classical
rw [vectorToPolynomial]
rw [← (radixTwoIndexEquiv k).sum_comp
(fun i => Polynomial.monomial i.1 (a i))]
rw [Fintype.sum_prod_type]
simp only [Fin.sum_univ_two]
simp [vectorToPolynomial, Polynomial.monomial_comp,
Finset.mul_sum, Finset.sum_add_distrib, radixTwoIndexEquiv_zero,
radixTwoIndexEquiv_one, radixTwoIndexEquiv_zero_eq,
radixTwoIndexEquiv_one_eq, evenCoeffs, oddCoeffs, evenIndex, oddIndex,
Polynomial.X_pow_eq_monomial, Polynomial.C_mul_monomial,
Polynomial.X_mul_monomial, pow_mul, Nat.mul_comm, Nat.add_comm]A DFT coordinate is polynomial evaluation at the corresponding root power.
private theorem dft_apply_eq_eval [CommSemiring K] {n : Nat}
(omega : K) (a : CoeffVector K n) (i : Fin n) :
dft omega a i = (vectorToPolynomial a).eval (omega ^ i.1) := by
rw [dft_eq_pointValues]
rflThe lower DFT half splits into even and odd child transforms.
private theorem dft_lower_split [CommRing K] {k : Nat}
(omega : K) (a : PowTwoVec K (k + 1)) (j : Fin (2 ^ k)) :
dft omega a (lowerHalfIndex j) =
dft (omega ^ 2) (evenCoeffs a) j +
omega ^ j.1 * dft (omega ^ 2) (oddCoeffs a) j := by
rw [dft_apply_eq_eval, lowerHalfIndex_val, polynomial_evenOdd_split]
rw [Polynomial.eval_add, Polynomial.eval_comp, Polynomial.eval_mul,
Polynomial.eval_X, Polynomial.eval_comp]
simp only [Polynomial.eval_pow, Polynomial.eval_X]
have hsquare : (omega ^ j.1) ^ 2 = (omega ^ 2) ^ j.1 := by
rw [← pow_mul, ← pow_mul]
congr 1
omega
rw [hsquare]
rw [← dft_apply_eq_eval, ← dft_apply_eq_eval]The upper DFT half uses the negative butterfly combination.
private theorem dft_upper_split [Field K] [CharZero K] {k : Nat}
{omega : K} (homega : IsPrimitiveRoot omega (2 ^ (k + 1)))
(a : PowTwoVec K (k + 1)) (j : Fin (2 ^ k)) :
dft omega a (upperHalfIndex j) =
dft (omega ^ 2) (evenCoeffs a) j -
omega ^ j.1 * dft (omega ^ 2) (oddCoeffs a) j := by
have horder : 2 ^ (k + 1) = 2 * 2 ^ k := by
simp [pow_succ, Nat.mul_comm]
have hhalf : omega ^ (2 ^ k) = -1 := by
apply primitiveRoot_half_pow_eq_neg_one (by positivity)
simpa [horder] using homega
have hfull : omega ^ (2 ^ (k + 1)) = 1 := homega.pow_eq_one
have hsquare :
(omega ^ (2 ^ k + j.1)) ^ 2 = (omega ^ 2) ^ j.1 := by
calc
(omega ^ (2 ^ k + j.1)) ^ 2 =
omega ^ ((2 ^ k + j.1) * 2) :=
(pow_mul omega (2 ^ k + j.1) 2).symm
_ = omega ^ (2 ^ (k + 1) + 2 * j.1) := by
congr 1
rw [pow_succ]
omega
_ = omega ^ (2 ^ (k + 1)) * omega ^ (2 * j.1) := by
rw [pow_add]
_ = (omega ^ 2) ^ j.1 := by
rw [hfull, one_mul]
rw [pow_mul]
have hpoint : omega ^ (2 ^ k + j.1) = -(omega ^ j.1) := by
rw [pow_add, hhalf]
ring
rw [dft_apply_eq_eval, upperHalfIndex_val, polynomial_evenOdd_split]
rw [Polynomial.eval_add, Polynomial.eval_comp, Polynomial.eval_mul,
Polynomial.eval_X, Polynomial.eval_comp]
simp only [Polynomial.eval_pow, Polynomial.eval_X]
rw [hsquare, hpoint]
rw [← dft_apply_eq_eval, ← dft_apply_eq_eval]
ringTwo successor-size vectors are equal when both contiguous halves agree.
private theorem powTwoVec_ext_halves {K : Type*} {k : Nat}
{f g : PowTwoVec K (k + 1)}
(hlower : ∀ j : Fin (2 ^ k), f (lowerHalfIndex j) = g (lowerHalfIndex j))
(hupper : ∀ j : Fin (2 ^ k), f (upperHalfIndex j) = g (upperHalfIndex j)) :
f = g := by
funext i
have h : ∀ t : Fin (2 ^ k + 2 ^ k),
f ((powTwoSuccEquiv k).symm t) = g ((powTwoSuccEquiv k).symm t) := by
intro t
refine Fin.addCases ?_ ?_ t
· intro j
exact hlower j
· intro j
exact hupper j
simpa using h (powTwoSuccEquiv k i)One butterfly layer combines the two half-size DFTs into the full DFT.
theorem butterflyLayer_dft [Field K] [CharZero K] {k : Nat}
{omega : K} (homega : IsPrimitiveRoot omega (2 ^ (k + 1)))
(a : PowTwoVec K (k + 1)) :
butterflyLayer omega
(dft (omega ^ 2) (evenCoeffs a))
(dft (omega ^ 2) (oddCoeffs a)) =
dft omega a := by
apply powTwoVec_ext_halves
· intro j
rw [butterflyLayer_lower, dft_lower_split]
· intro j
rw [butterflyLayer_upper, dft_upper_split homega]The actual recursive radix-2 execution computes the generic DFT.
theorem recursiveFFT_eq_dft [Field K] [CharZero K] {k : Nat}
{omega : K} (homega : IsPrimitiveRoot omega (2 ^ k))
(a : PowTwoVec K k) :
recursiveFFT omega a = dft omega a := by
induction k generalizing omega with
| zero =>
funext i
fin_cases i
simp [recursiveFFT, recursiveFFTExec, dft]
| succ k ih =>
have hsquare : IsPrimitiveRoot (omega ^ 2) (2 ^ k) := by
apply primitiveRoot_square (by positivity)
simpa [pow_succ, Nat.mul_comm] using homega
have hchildRoot :
twiddleChildRoot k omega
(twiddlePowersAuxExec omega (2 ^ k) 1) = omega ^ 2 := by
by_cases hk : k = 0
· subst k
simp only [twiddleChildRoot, if_pos rfl]
exact homega.pow_eq_one.symm
· exact twiddleChildRoot_eq_square (Nat.pos_of_ne_zero hk) omega
simp only [recursiveFFT, recursiveFFTExec]
rw [hchildRoot]
change butterflyLayer omega
(recursiveFFT (omega ^ 2) (evenCoeffs a))
(recursiveFFT (omega ^ 2) (oddCoeffs a)) = dft omega a
rw [ih hsquare, ih hsquare]
exact butterflyLayer_dft homega aThe recursive inverse agrees with the algebraic inverse DFT.
theorem recursiveIFFT_eq_idft [Field K] [CharZero K] {k : Nat}
{omega : K} (homega : IsPrimitiveRoot omega (2 ^ k))
(a : PowTwoVec K k) :
recursiveIFFT omega a = idft omega a := by
funext i
simp [recursiveIFFT, idft,
recursiveFFT_eq_dft (primitiveRoot_inv homega)]Recursive inverse after recursive forward transform is the identity.
theorem recursiveIFFT_recursiveFFT [Field K] [CharZero K] {k : Nat}
{omega : K} (homega : IsPrimitiveRoot omega (2 ^ k))
(a : PowTwoVec K k) :
recursiveIFFT omega (recursiveFFT omega a) = a := by
rw [recursiveIFFT_eq_idft homega, recursiveFFT_eq_dft homega,
idft_dft (by positivity) homega]Recursive forward transform after recursive inverse is the identity.
theorem recursiveFFT_recursiveIFFT [Field K] [CharZero K] {k : Nat}
{omega : K} (homega : IsPrimitiveRoot omega (2 ^ k))
(a : PowTwoVec K k) :
recursiveFFT omega (recursiveIFFT omega a) = a := by
rw [recursiveFFT_eq_dft homega, recursiveIFFT_eq_idft homega,
dft_idft (by positivity) homega]end Chapter30end CLRSCLRSLean.FourthEdition.Chapter_30.Section_30_2_DFT_And_FFT.RecursiveFFT.Costs
Chapter 30.2: Recursive FFT costs and padding
All exact counts below are projections of recursiveFFTExec in the
shared butterfly arithmetic model. Each product used by both output slots is
charged once; this does not count duplicated evaluation of function-valued
outputs. The independent
numeric closed form is introduced only after those execution equations.
namespace CLRSnamespace Chapter30Work measured by the canonical value-producing FFT execution.
def recursiveFFTWork [Ring K] {k : Nat} (omega : K) (a : PowTwoVec K k) : Nat :=
(recursiveFFTExec omega a).workThe execution charges exactly one lower sum and one upper subtraction per output at each recursive level.
theorem recursiveFFTExec_addSubtractions [Ring K] {k : Nat}
(omega : K) (a : PowTwoVec K k) :
(recursiveFFTExec omega a).addSubtractions = k * 2 ^ k := by
induction k generalizing omega with
| zero => rfl
| succ k ih =>
simp only [recursiveFFTExec]
rw [ih, ih]
simp [butterflyLayerFromTwiddleExec, pow_succ]
ringThe execution charges every butterfly product and every successive twiddle update, giving the same exact count as addition/subtraction.
theorem recursiveFFTExec_multiplications [Ring K] {k : Nat}
(omega : K) (a : PowTwoVec K k) :
(recursiveFFTExec omega a).multiplications = k * 2 ^ k := by
induction k generalizing omega with
| zero => rfl
| succ k ih =>
simp only [recursiveFFTExec]
rw [ih, ih]
simp [butterflyLayerFromTwiddleExec,
twiddlePowersAuxExec_multiplications, pow_succ]
ringExact charged work of recursive FFT under the shared-product model.
theorem recursiveFFTWork_exact [Ring K] {k : Nat}
(omega : K) (a : PowTwoVec K k) :
recursiveFFTWork omega a = 2 * k * 2 ^ k := by
rw [recursiveFFTWork, FFTExecution.work,
recursiveFFTExec_addSubtractions, recursiveFFTExec_multiplications]
ringNumeric closed form extracted from the execution theorem.
def radix2FFTWork (k : Nat) : Nat := 2 * k * 2 ^ kThe execution-derived work agrees with the public numeric cost function.
theorem recursiveFFTWork_eq_radix2FFTWork [Ring K] {k : Nat}
(omega : K) (a : PowTwoVec K k) :
recursiveFFTWork omega a = radix2FFTWork k := by
rw [recursiveFFTWork_exact]
rfl
The exact-power execution work is Theta(k * 2^k).
theorem radix2FFTWork_bigTheta :
Chapter03.isBigTheta
(fun k => (radix2FFTWork k : ℝ))
(fun k => (k : ℝ) * (2 : ℝ) ^ k) := by
constructor
· refine (Chapter03.isBigO_iff _ _).mpr ⟨2, by norm_num, 0, ?_⟩
intro k _
rw [abs_of_nonneg (by positivity), abs_of_nonneg (by positivity)]
norm_num [radix2FFTWork]
simpa only [mul_assoc] using
(le_refl ((2 : ℝ) * (k : ℝ) * (2 : ℝ) ^ k))
· refine (Chapter03.isBigOmega_iff _ _).mpr ⟨(2 : ℝ), by norm_num, 0, ?_⟩
intro k _
rw [abs_of_nonneg (by positivity), abs_of_nonneg (by positivity)]
norm_num [radix2FFTWork]
simpa only [mul_assoc] using
(le_refl ((2 : ℝ) * (k : ℝ) * (2 : ℝ) ^ k))Least power-of-two exponent used for total zero padding.
def fftExponent (n : Nat) : Nat := Nat.clog 2 (max 1 n)
Nonempty power-of-two transform capacity covering n.
def fftCapacity (n : Nat) : Nat := 2 ^ fftExponent nWork of the canonical recursive FFT at the padded capacity.
def paddedFFTWork (n : Nat) : Nat := radix2FFTWork (fftExponent n)Padding never loses an original slot.
theorem fftCapacity_ge (n : Nat) : n ≤ fftCapacity n := by
exact (Nat.le_max_right 1 n).trans
(Nat.le_pow_clog (by norm_num) (max 1 n))Every FFT capacity is explicitly a power of two.
theorem fftCapacity_isPowerOfTwo (n : Nat) : ∃ k, fftCapacity n = 2 ^ k :=
⟨fftExponent n, rfl⟩FFT capacity is always positive, including at input zero.
theorem fftCapacity_pos (n : Nat) : 0 < fftCapacity n := by
simp [fftCapacity]Above one, padding uses strictly less than twice the input capacity.
theorem fftCapacity_lt_two_mul {n : Nat} (hn : 1 < n) :
fftCapacity n < 2 * n := by
have hpred := Nat.pow_pred_clog_lt_self (b := 2) (by norm_num) hn
have hclog : 0 < Nat.clog 2 n := Nat.clog_pos (by norm_num) hn
rw [fftCapacity, fftExponent, max_eq_right (by omega)]
rw [show Nat.clog 2 n = (Nat.clog 2 n).pred + 1 from
(Nat.succ_pred_eq_of_pos hclog).symm, pow_succ]
rw [Nat.mul_comm]
exact (Nat.mul_lt_mul_left (a := 2) (b := 2 ^ (Nat.clog 2 n).pred)
(c := n) (by norm_num)).mpr hpredThe selected padding exponent is monotone in the requested capacity.
theorem fftExponent_monotone : Monotone fftExponent := by
intro m n hmn
exact Nat.clog_mono_right 2 (max_le_max_left 1 hmn)The selected power-of-two capacity is monotone.
theorem fftCapacity_monotone : Monotone fftCapacity := by
intro m n hmn
exact Nat.pow_le_pow_right (by norm_num) (fftExponent_monotone hmn)Exact radix-2 FFT work is monotone in the exponent.
private theorem radix2FFTWork_monotone : Monotone radix2FFTWork := by
intro k l hkl
unfold radix2FFTWork
have hpow : 2 ^ k ≤ 2 ^ l := Nat.pow_le_pow_right (by norm_num) hkl
simpa [Nat.mul_assoc] using Nat.mul_le_mul_left 2 (Nat.mul_le_mul hkl hpow)Padded FFT work is monotone in the advertised input capacity.
theorem paddedFFTWork_monotone : Monotone paddedFFTWork := by
intro m n hmn
exact radix2FFTWork_monotone (fftExponent_monotone hmn)Padding is exact on power-of-two input sizes.
theorem paddedFFTWork_pow (k : Nat) :
paddedFFTWork (2 ^ k) = 2 * k * 2 ^ k := by
rw [paddedFFTWork, fftExponent, max_eq_right Nat.one_le_two_pow,
Nat.clog_pow 2 k (by norm_num)]
rflOn exact powers of two, padded work has the critical linear-log scale.
private theorem paddedFFTWork_exactPower_bigTheta :
Chapter03.isBigTheta
(fun k : Nat => (paddedFFTWork (2 ^ k) : ℝ))
(fun k : Nat => ((k : ℝ) + 1) * (2 : ℝ) ^ k) := by
constructor
· refine (Chapter03.isBigO_iff _ _).mpr ⟨2, by norm_num, 0, ?_⟩
intro k _
rw [abs_of_nonneg (Nat.cast_nonneg _), abs_of_nonneg (by positivity)]
rw [paddedFFTWork_pow]
push_cast
have hk : 0 ≤ (k : ℝ) := Nat.cast_nonneg k
have hpow : 0 ≤ (2 : ℝ) ^ k := by positivity
nlinarith
· refine (Chapter03.isBigOmega_iff _ _).mpr ⟨1, by norm_num, 1, ?_⟩
intro k hk_one
rw [abs_of_nonneg (by positivity), abs_of_nonneg (Nat.cast_nonneg _)]
rw [paddedFFTWork_pow]
push_cast
have hk : 1 ≤ (k : ℝ) := by exact_mod_cast hk_one
have hpow : 0 ≤ (2 : ℝ) ^ k := by positivity
nlinarith
Zero padding lifts the exact power-of-two count to the textbook
Theta(n log n) scale on every input size.
theorem paddedFFTWork_allInput_bigTheta :
Chapter03.isBigTheta
(fun n : Nat => (paddedFFTWork n : ℝ))
(Chapter04.realLogLogScale 2 2) := by
have hcritical :
Chapter03.isBigTheta
(fun n : Nat => (paddedFFTWork n : ℝ))
(Chapter04.criticalPowerLogScale 2 2) :=
Chapter04.allInput_bigTheta_of_criticalPowerLogScale 2 2
(fun n : Nat => (paddedFFTWork n : ℝ)) (by norm_num) (by norm_num)
(Chapter04.monotoneAbs_natCast paddedFFTWork_monotone)
paddedFFTWork_exactPower_bigTheta
exact Chapter03.isBigTheta_trans hcritical
(Chapter04.criticalPowerLogScale_isBigTheta_realLogLogScale 2 2
(by norm_num) (by norm_num))Zero-pad a vector to its least nonempty power-of-two capacity.
def zeroPadToFFTCapacity [Zero K] {n : Nat} (a : CoeffVector K n) :
CoeffVector K (fftCapacity n) :=
fun i => if h : i.1 < n then a ⟨i.1, h⟩ else 0
The original coefficient at i survives zero padding.
theorem zeroPadToFFTCapacity_original [Zero K] {n : Nat}
(a : CoeffVector K n) (i : Fin n) :
zeroPadToFFTCapacity a ⟨i.1, i.2.trans_le (fftCapacity_ge n)⟩ = a i := by
simp [zeroPadToFFTCapacity]Every added padding slot is zero.
theorem zeroPadToFFTCapacity_added [Zero K] {n : Nat}
(a : CoeffVector K n) (i : Fin (fftCapacity n)) (hi : n ≤ i.1) :
zeroPadToFFTCapacity a i = 0 := by
simp [zeroPadToFFTCapacity, Nat.not_lt.mpr hi]Padded work remains attached to the actual recursive execution.
theorem recursiveFFTExec_zeroPad_work [Ring K] {n : Nat}
(omega : K) (a : CoeffVector K n) :
(recursiveFFTExec (k := fftExponent n) omega
(zeroPadToFFTCapacity a)).work = paddedFFTWork n := by
rw [← recursiveFFTWork]
rw [recursiveFFTWork_eq_radix2FFTWork]
rflend Chapter30end CLRSCLRSLean.FourthEdition.Chapter_30.Section_30_2_DFT_And_FFT.PolynomialMultiplication
Chapter 30.2: FFT polynomial multiplication
The generic pipeline below composes three actual recursive FFT executions, one pointwise-product execution, and one inverse-scaling execution. Its counters charge field additions/subtractions and multiplications. Coefficient reads, finite-index bookkeeping, primitive-root certificates, and polynomial reconstruction are outside this arithmetic-operation metric.
namespace CLRSnamespace Chapter30open PolynomialPolynomial result and arithmetic counters of the FFT multiplication pipeline.
structure FFTMultiplicationExecution (K : Type*) [Semiring K] where
value : Polynomial K
addSubtractions : Nat
multiplications : NatTotal charged field arithmetic of an FFT multiplication execution.
def FFTMultiplicationExecution.work [Semiring K]
(r : FFTMultiplicationExecution K) : Nat :=
r.addSubtractions + r.multiplicationsPointwise scalar multiplication, used for inverse-transform scaling.
def scaleVectorExec [Mul K] {n : Nat} (c : K) (a : CoeffVector K n) :
VectorArithmeticExecution K n :=
⟨fun i => c * a i, 0, n⟩Inverse scaling charges exactly one multiplication per output slot.
theorem scaleVectorExec_work_exact [Mul K] {n : Nat}
(c : K) (a : CoeffVector K n) :
(scaleVectorExec c a).work = n := by
simp [scaleVectorExec, VectorArithmeticExecution.work]The canonical fixed-capacity FFT multiplication execution.
noncomputable def fftMultiplyExecAt [Field K] {k : Nat} (omega : K)
(p q : K[X]) : FFTMultiplicationExecution K :=
let a := coeffVector (2 ^ k) p
let b := coeffVector (2 ^ k) q
let leftRun := recursiveFFTExec omega a
let rightRun := recursiveFFTExec omega b
let productRun := pointwiseMulExec leftRun.value rightRun.value
let inverseRun := recursiveFFTExec omega⁻¹ productRun.value
let scaleRun := scaleVectorExec ((2 ^ k : Nat) : K)⁻¹ inverseRun.value
⟨vectorToPolynomial scaleRun.value,
leftRun.addSubtractions + rightRun.addSubtractions +
inverseRun.addSubtractions + productRun.additions + scaleRun.additions,
leftRun.multiplications + rightRun.multiplications +
inverseRun.multiplications + productRun.multiplications +
scaleRun.multiplications⟩Value projection of fixed-capacity FFT multiplication.
noncomputable def fftMultiplyAt [Field K] {k : Nat} (omega : K)
(p q : K[X]) : K[X] :=
(fftMultiplyExecAt (k := k) omega p q).valuePublic erasure equation for the costed multiplication execution.
theorem fftMultiplyExecAt_value [Field K] {k : Nat} (omega : K)
(p q : K[X]) :
(fftMultiplyExecAt (k := k) omega p q).value =
fftMultiplyAt (k := k) omega p q := rflThe vector produced by the actual execution is the semantic inverse-DFT pipeline. This isolates all counter fields from polynomial correctness.
theorem fftMultiplyExecAt_vector_semantics [Field K] [CharZero K] {k : Nat}
{omega : K} (homega : IsPrimitiveRoot omega (2 ^ k)) (p q : K[X]) :
let a := coeffVector (2 ^ k) p
let b := coeffVector (2 ^ k) q
let leftRun := recursiveFFTExec omega a
let rightRun := recursiveFFTExec omega b
let productRun := pointwiseMulExec leftRun.value rightRun.value
let inverseRun := recursiveFFTExec omega⁻¹ productRun.value
let scaleRun := scaleVectorExec ((2 ^ k : Nat) : K)⁻¹ inverseRun.value
scaleRun.value =
idft omega (pointwiseMul (dft omega a) (dft omega b)) := by
dsimp only
change recursiveIFFT omega
(pointwiseMul (recursiveFFT omega (coeffVector (2 ^ k) p))
(recursiveFFT omega (coeffVector (2 ^ k) q))) =
idft omega
(pointwiseMul (dft omega (coeffVector (2 ^ k) p))
(dft omega (coeffVector (2 ^ k) q)))
rw [recursiveIFFT_eq_idft homega,
recursiveFFT_eq_dft homega, recursiveFFT_eq_dft homega]Fixed-capacity FFT multiplication is exact whenever the product fits in the declared transform capacity.
theorem fftMultiplyAt_correct [Field K] [CharZero K] {k : Nat}
{omega : K} (homega : IsPrimitiveRoot omega (2 ^ k))
(p q : K[X])
(hfit : (p * q).degree < ((2 ^ k : Nat) : WithBot Nat)) :
fftMultiplyAt (k := k) omega p q = p * q := by
rw [fftMultiplyAt, fftMultiplyExecAt]
dsimp only
rw [fftMultiplyExecAt_vector_semantics homega p q]
rw [idft_pointwiseMul (by positivity) homega]
calc
vectorToPolynomial
(cyclicConvolution (by positivity)
(coeffVector (2 ^ k) p) (coeffVector (2 ^ k) q)) =
vectorToPolynomial (coeffVector (2 ^ k) (p * q)) := by
congr 1
exact cyclicConvolution_eq_coeffVector_mul
(by positivity) p q hfit
_ = p * q := vectorToPolynomial_coeffVector (p * q) hfitA support-bound wrapper deriving the minimal no-wrap premise from operand degree bounds and a declared capacity.
theorem fftMultiplyAt_correct_of_degree_lt [Field K] [CharZero K]
{k m n : Nat} {omega : K} (homega : IsPrimitiveRoot omega (2 ^ k))
(p q : K[X]) (hp : p.degree < m) (hq : q.degree < n)
(hcapacity : m + n ≤ 2 ^ k) :
fftMultiplyAt (k := k) omega p q = p * q := by
apply fftMultiplyAt_correct homega p q
by_cases hp0 : p = 0
· subst p
simp only [zero_mul, Polynomial.degree_zero]
exact WithBot.bot_lt_coe _
by_cases hq0 : q = 0
· subst q
simp only [mul_zero, Polynomial.degree_zero]
exact WithBot.bot_lt_coe _
have hpnat : p.natDegree < m :=
(Polynomial.natDegree_lt_iff_degree_lt hp0).mpr hp
have hqnat : q.natDegree < n :=
(Polynomial.natDegree_lt_iff_degree_lt hq0).mpr hq
have hpq0 : p * q ≠ 0 := mul_ne_zero hp0 hq0
apply (Polynomial.natDegree_lt_iff_degree_lt hpq0).mp
rw [Polynomial.natDegree_mul hp0 hq0]
omegaArbitrary-input complex wrapper
Positive coefficient capacity associated with a polynomial, including the zero polynomial.
def polySize [Semiring K] (p : K[X]) : Nat :=
p.natDegree + 1The coefficient-capacity convention assigns every polynomial positive size.
A symmetric positive size sufficient for multiplying two operands.
The symmetric advertised size for a multiplication is positive.
theorem multiplicationInputSize_pos [Semiring K] (p q : K[X]) :
0 < multiplicationInputSize p q := by
have hmax : 0 < max (polySize p) (polySize q) :=
(polySize_pos p).trans_le (Nat.le_max_left _ _)
simp [multiplicationInputSize, hmax]Every polynomial fits in its positive coefficient-size convention.
theorem degree_lt_polySize [Semiring K] (p : K[X]) :
p.degree < ((polySize p : Nat) : WithBot Nat) := by
by_cases hp : p = 0
· subst p
simp [polySize]
· apply (Polynomial.natDegree_lt_iff_degree_lt hp).mp
simp [polySize]The symmetric wrapper size strictly exceeds the product degree, including zero and constant operands.
theorem mul_degree_lt_multiplicationInputSize [Semiring K] (p q : K[X]) :
(p * q).degree <
((multiplicationInputSize p q : Nat) : WithBot Nat) := by
have hnat :
(p * q).natDegree < multiplicationInputSize p q := by
calc
(p * q).natDegree ≤ p.natDegree + q.natDegree :=
Polynomial.natDegree_mul_le
_ < 2 * max (polySize p) (polySize q) := by
have hpmax : p.natDegree + 1 ≤ max (polySize p) (polySize q) := by
simpa [polySize] using Nat.le_max_left (polySize p) (polySize q)
have hqmax : q.natDegree + 1 ≤ max (polySize p) (polySize q) := by
simpa [polySize] using Nat.le_max_right (polySize p) (polySize q)
omega
_ = multiplicationInputSize p q := rfl
exact (Polynomial.degree_le_natDegree.trans_lt (by exact_mod_cast hnat))Least radix-2 exponent selected by the complex convenience wrapper.
def complexFFTExponent (p q : ℂ[X]) : Nat :=
fftExponent (multiplicationInputSize p q)Power-of-two transform capacity selected by the complex wrapper.
def complexFFTCapacity (p q : ℂ[X]) : Nat :=
2 ^ complexFFTExponent p qThe automatically selected complex transform capacity is positive.
theorem complexFFTCapacity_pos (p q : ℂ[X]) :
0 < complexFFTCapacity p q := by
simp [complexFFTCapacity]The automatically chosen complex capacity contains the entire product.
theorem complex_product_fits (p q : ℂ[X]) :
(p * q).degree <
((complexFFTCapacity p q : Nat) : WithBot Nat) := by
exact (mul_degree_lt_multiplicationInputSize p q).trans_le
(by
exact_mod_cast (fftCapacity_ge (multiplicationInputSize p q)))Positive-sign principal root matching the generic DFT convention.
noncomputable def complexFFTRoot (p q : ℂ[X]) : ℂ :=
Complex.exp
(2 * Real.pi * Complex.I / (complexFFTCapacity p q : ℂ))The complex wrapper's principal root is primitive at its selected capacity.
theorem complexFFTRoot_isPrimitive (p q : ℂ[X]) :
IsPrimitiveRoot (complexFFTRoot p q) (complexFFTCapacity p q) := by
simpa [complexFFTRoot] using
Complex.isPrimitiveRoot_exp (complexFFTCapacity p q)
(Nat.ne_of_gt (complexFFTCapacity_pos p q))Arbitrary-input exact complex polynomial multiplication.
noncomputable def complexFFTMultiply (p q : ℂ[X]) : ℂ[X] :=
fftMultiplyAt (k := complexFFTExponent p q) (complexFFTRoot p q) p qThe complex wrapper constructs a sufficient capacity and primitive root internally, so callers need no fit premise.
theorem complexFFTMultiply_correct (p q : ℂ[X]) :
complexFFTMultiply p q = p * q := by
apply fftMultiplyAt_correct (complexFFTRoot_isPrimitive p q)
exact complex_product_fits p qExecution-attached multiplication costs
Exact-capacity composition: two forward FFTs, one inverse-root FFT, one pointwise product per slot, and one inverse-scale product per slot.
def radix2FFTMultiplyWork (k : Nat) : Nat :=
3 * radix2FFTWork k + 2 * 2 ^ kThe work field of the actual multiplication execution equals the numeric composition; no pipeline stage is charged by a detached recurrence.
theorem fftMultiplyExecAt_work_exact [Field K] {k : Nat}
(omega : K) (p q : K[X]) :
(fftMultiplyExecAt (k := k) omega p q).work =
radix2FFTMultiplyWork k := by
simp [fftMultiplyExecAt, FFTMultiplicationExecution.work,
radix2FFTMultiplyWork, recursiveFFTExec_addSubtractions,
recursiveFFTExec_multiplications, pointwiseMulExec, scaleVectorExec,
radix2FFTWork]
ring
Transform exponent charged for two operands each advertised with
coefficient capacity n.
def fftMultiplyExponent (n : Nat) : Nat :=
fftExponent (2 * max 1 n)
Declared all-input arithmetic cost. Every operand with degree < n is
charged at this capacity, even if its leading coefficients vanish.
def fftMultiplyWork (n : Nat) : Nat :=
radix2FFTMultiplyWork (fftMultiplyExponent n)
The actual execution selected for advertised operand capacity n.
noncomputable def fftMultiplyExecution [Field K] (n : Nat) (omega : K)
(p q : K[X]) : FFTMultiplicationExecution K :=
fftMultiplyExecAt (k := fftMultiplyExponent n) omega p qThe declared cost is exactly the work field of the selected execution.
theorem fftMultiplyExecution_work_eq [Field K] (n : Nat) (omega : K)
(p q : K[X]) :
(fftMultiplyExecution n omega p q).work = fftMultiplyWork n := by
exact fftMultiplyExecAt_work_exact omega p qCorrectness and cost use the same selected execution for bounded operands.
theorem fftMultiplyExecution_correct [Field K] [CharZero K] {n : Nat}
{omega : K}
(homega : IsPrimitiveRoot omega (2 ^ fftMultiplyExponent n))
(p q : K[X]) (hp : p.degree < n) (hq : q.degree < n) :
(fftMultiplyExecution n omega p q).value = p * q := by
rw [fftMultiplyExecution, fftMultiplyExecAt_value]
apply fftMultiplyAt_correct_of_degree_lt (m := n) (n := n) homega p q hp hq
have hnmax : n ≤ max 1 n := Nat.le_max_right 1 n
exact (by omega : n + n ≤ 2 * max 1 n) |>.trans
(fftCapacity_ge (2 * max 1 n))Exact-capacity FFT multiplication work is monotone in the radix exponent.
private theorem radix2FFTMultiplyWork_monotone :
Monotone radix2FFTMultiplyWork := by
intro k l hkl
have hpow : 2 ^ k ≤ 2 ^ l := Nat.pow_le_pow_right (by norm_num) hkl
have hfft : radix2FFTWork k ≤ radix2FFTWork l := by
unfold radix2FFTWork
simpa [Nat.mul_assoc] using
Nat.mul_le_mul_left 2 (Nat.mul_le_mul hkl hpow)
exact Nat.add_le_add
(Nat.mul_le_mul_left 3 hfft) (Nat.mul_le_mul_left 2 hpow)The multiplication transform exponent is monotone in operand capacity.
theorem fftMultiplyExponent_monotone : Monotone fftMultiplyExponent := by
intro m n hmn
exact fftExponent_monotone
(Nat.mul_le_mul_left 2 (max_le_max_left 1 hmn))The advertised all-input multiplication work is monotone.
theorem fftMultiplyWork_monotone : Monotone fftMultiplyWork := by
intro m n hmn
exact radix2FFTMultiplyWork_monotone
(fftMultiplyExponent_monotone hmn)Power-of-two operand capacity selects the next radix exponent.
theorem fftMultiplyExponent_pow (k : Nat) :
fftMultiplyExponent (2 ^ k) = k + 1 := by
rw [fftMultiplyExponent, max_eq_right Nat.one_le_two_pow]
rw [show 2 * 2 ^ k = 2 ^ (k + 1) by
simp [pow_succ, Nat.mul_comm]]
rw [fftExponent, max_eq_right Nat.one_le_two_pow,
Nat.clog_pow 2 (k + 1) (by norm_num)]Closed form of multiplication work at advertised power-of-two operand capacity. The linear pointwise/scaling term is retained explicitly.
theorem fftMultiplyWork_pow (k : Nat) :
fftMultiplyWork (2 ^ k) = (12 * k + 16) * 2 ^ k := by
rw [fftMultiplyWork, fftMultiplyExponent_pow]
simp [radix2FFTMultiplyWork, radix2FFTWork, pow_succ]
ringExact-power multiplication work has the critical linear-log scale.
private theorem fftMultiplyWork_exactPower_bigTheta :
Chapter03.isBigTheta
(fun k : Nat => (fftMultiplyWork (2 ^ k) : ℝ))
(fun k : Nat => ((k : ℝ) + 1) * (2 : ℝ) ^ k) := by
constructor
· refine (Chapter03.isBigO_iff _ _).mpr ⟨16, by norm_num, 0, ?_⟩
intro k _
rw [abs_of_nonneg (Nat.cast_nonneg _), abs_of_nonneg (by positivity)]
rw [fftMultiplyWork_pow]
push_cast
have hk : 0 ≤ (k : ℝ) := Nat.cast_nonneg k
have hpow : 0 ≤ (2 : ℝ) ^ k := by positivity
have hkp : 0 ≤ (k : ℝ) * (2 : ℝ) ^ k := mul_nonneg hk hpow
nlinarith
· refine (Chapter03.isBigOmega_iff _ _).mpr ⟨1, by norm_num, 0, ?_⟩
intro k _
rw [abs_of_nonneg (by positivity), abs_of_nonneg (Nat.cast_nonneg _)]
rw [fftMultiplyWork_pow]
push_cast
have hk : 0 ≤ (k : ℝ) := Nat.cast_nonneg k
have hpow : 0 ≤ (2 : ℝ) ^ k := by positivity
have hkp : 0 ≤ (k : ℝ) * (2 : ℝ) ^ k := mul_nonneg hk hpow
nlinarith
Fixed-capacity FFT multiplication has all-input Theta(n log n) charged
field arithmetic.
theorem fftMultiplyWork_allInput_bigTheta :
Chapter03.isBigTheta
(fun n : Nat => (fftMultiplyWork n : ℝ))
(Chapter04.realLogLogScale 2 2) := by
have hcritical :
Chapter03.isBigTheta
(fun n : Nat => (fftMultiplyWork n : ℝ))
(Chapter04.criticalPowerLogScale 2 2) :=
Chapter04.allInput_bigTheta_of_criticalPowerLogScale 2 2
(fun n : Nat => (fftMultiplyWork n : ℝ)) (by norm_num) (by norm_num)
(Chapter04.monotoneAbs_natCast fftMultiplyWork_monotone)
fftMultiplyWork_exactPower_bigTheta
exact Chapter03.isBigTheta_trans hcritical
(Chapter04.criticalPowerLogScale_isBigTheta_realLogLogScale 2 2
(by norm_num) (by norm_num))end Chapter30end CLRSImports
30.3. Efficient FFT Implementations
The total bit-reversal permutation exposes structural even/odd laws,
fixed-width bit semantics, involution, and an exact one-move-per-element count.
The functional iterative FFT folds globally ordered stages over that
permutation. Its stage-prefix invariant factors every nonfinal stage over the
two contiguous halves, proving equality with the recursive FFT and the generic
DFT. The execution records exactly one bit-reversal move per input,
k * 2^k additions/subtractions, and k * 2^k multiplications; padding lifts
the total work to the all-input Theta(n log n) scale.
The explicit layered network stores recursive stage circuits whose leaves are
actual butterflies with fixed twiddle constants. Evaluating those stored
gates is proved equal to the iterative FFT, and the same circuit syntax has
k * 2^(k-1) butterflies and butterfly depth k; expanding a butterfly to
one multiplication and two addition/subtraction gates gives
3 * k * 2^(k-1) primitive gates and primitive depth 2 * k. Twiddle powers
are circuit constants and bit reversal is wiring in this circuit model, while
the functional execution separately charges successive twiddle updates.
Mutable arrays, aliasing and in-place loop semantics; RAM, cache, allocator, SIMD, GPU, and communication costs; parallel scheduling and processor bounds; floating-point error and numerical stability; number-theoretic-transform specialization; code generation; exercises; and Problems 30-1 through 30-6 remain outside this exact functional boundary.
Implementation pages:
namespace CLRSnamespace Chapter30end Chapter30end CLRSDefinitions and proofs
CLRSLean.FourthEdition.Chapter_30.Section_30_3_Efficient_FFT_Implementations.BitReversal
Chapter 30.3: Bit-reversal permutation
The executable copy and its index equivalence expose the bit-reversal step of the iterative radix-2 FFT without introducing partial array indexing.
namespace CLRSnamespace Chapter30Decompose a successor-size index into its quotient and low bit.
private def powTwoSuccToQuotBit (k : Nat) :
Fin (2 ^ (k + 1)) ≃ Fin (2 ^ k) × Fin 2 :=
(finCongr (by simp [pow_succ])).trans
(finProdFinEquiv (m := 2 ^ k) (n := 2)).symmReassemble a leading bit and remaining fixed-width index.
private def bitRestToPowTwoSucc (k : Nat) :
Fin 2 × Fin (2 ^ k) ≃ Fin (2 ^ (k + 1)) :=
(finProdFinEquiv (m := 2) (n := 2 ^ k)).trans
(finCongr (by simp [pow_succ, Nat.mul_comm]))Reverse the declared low bits of a power-of-two index.
def bitReverseEquiv : (k : Nat) → Fin (2 ^ k) ≃ Fin (2 ^ k)
| 0 => Equiv.refl _
| k + 1 =>
(powTwoSuccToQuotBit k).trans
(((bitReverseEquiv k).prodCongr (Equiv.refl (Fin 2))).trans
((Equiv.prodComm (Fin (2 ^ k)) (Fin 2)).trans
(bitRestToPowTwoSucc k)))Reversing an even index places the recursively reversed index below.
@[simp] theorem bitReverseEquiv_even {k : Nat} (i : Fin (2 ^ k)) :
bitReverseEquiv (k + 1) (evenIndex i) =
lowerHalfIndex (bitReverseEquiv k i) := by
apply Fin.ext
simp [bitReverseEquiv, powTwoSuccToQuotBit, bitRestToPowTwoSucc,
evenIndex, lowerHalfIndex, powTwoSuccEquiv, Fin.divNat, Fin.modNat]Reversing an odd index places the recursively reversed index above.
@[simp] theorem bitReverseEquiv_odd {k : Nat} (i : Fin (2 ^ k)) :
bitReverseEquiv (k + 1) (oddIndex i) =
upperHalfIndex (bitReverseEquiv k i) := by
have hdiv : (2 * i.1 + 1) / 2 = i.1 := by omega
apply Fin.ext
simp [bitReverseEquiv, powTwoSuccToQuotBit, bitRestToPowTwoSucc,
oddIndex, upperHalfIndex, powTwoSuccEquiv, Fin.divNat, Fin.modNat, hdiv]Reversal mirrors every bit position inside the declared fixed width.
theorem bitReverseEquiv_testBit {k : Nat} (i : Fin (2 ^ k))
(j : Nat) (hj : j < k) :
Nat.testBit (bitReverseEquiv k i).1 j =
Nat.testBit i.1 (k - 1 - j) := by
induction k generalizing j with
| zero => omega
| succ k ih =>
obtain ⟨q, hq | hq⟩ := Nat.even_or_odd' i.1
· have hq_lt : q < 2 ^ k := by
have hi := i.2
simp [pow_succ] at hi
omega
let qi : Fin (2 ^ k) := ⟨q, hq_lt⟩
have hi : i = evenIndex qi := Fin.ext hq
subst i
by_cases htop : j = k
· subst j
rw [bitReverseEquiv_even]
simp only [lowerHalfIndex_val, evenIndex_val]
rw [show k + 1 - 1 - k = 0 by omega]
rw [Nat.testBit_lt_two_pow (bitReverseEquiv k qi).2]
simp [Nat.testBit_zero]
· have hjk : j < k := by omega
rw [bitReverseEquiv_even]
simp only [lowerHalfIndex_val, evenIndex_val]
rw [ih qi j hjk]
rw [show k + 1 - 1 - j = (k - 1 - j) + 1 by omega]
rw [Nat.testBit_add_one]
congr 1
omega
· have hq_lt : q < 2 ^ k := by
have hi := i.2
simp [pow_succ] at hi
omega
let qi : Fin (2 ^ k) := ⟨q, hq_lt⟩
have hi : i = oddIndex qi := Fin.ext hq
subst i
by_cases htop : j = k
· subst j
rw [bitReverseEquiv_odd]
simp only [upperHalfIndex_val, oddIndex_val]
rw [show k + 1 - 1 - k = 0 by omega]
rw [Nat.testBit_two_pow_add_eq]
rw [Nat.testBit_lt_two_pow (bitReverseEquiv k qi).2]
simp [Nat.testBit_zero]
· have hjk : j < k := by omega
rw [bitReverseEquiv_odd]
simp only [upperHalfIndex_val, oddIndex_val]
rw [Nat.testBit_two_pow_add_gt hjk]
rw [ih qi j hjk]
rw [show k + 1 - 1 - j = (k - 1 - j) + 1 by omega]
rw [Nat.testBit_add_one]
congr 1
omegaReversing the same fixed-width index twice is the identity.
theorem bitReverseEquiv_involutive (k : Nat) :
Function.Involutive (bitReverseEquiv k) := by
intro i
apply Fin.ext
apply Nat.eq_of_testBit_eq
intro j
by_cases hj : j < k
· rw [bitReverseEquiv_testBit _ j hj]
have hmirror : k - 1 - j < k := by omega
rw [bitReverseEquiv_testBit _ _ hmirror]
congr 1
omega
· have hkj : k ≤ j := Nat.le_of_not_gt hj
have hpow : 2 ^ k ≤ 2 ^ j := Nat.pow_le_pow_right (by omega) hkj
rw [Nat.testBit_lt_two_pow
((bitReverseEquiv k (bitReverseEquiv k i)).2.trans_le hpow)]
rw [Nat.testBit_lt_two_pow (i.2.trans_le hpow)]Result and data-movement counter for bit-reversal copy.
structure BitReverseExecution (K : Type*) (k : Nat) where
value : PowTwoVec K k
moves : NatCopy a vector into bit-reversed order by recursively grouping its even and odd coefficients. Every singleton leaf contributes one output move.
def bitReverseExec {K : Type*} :
{k : Nat} → PowTwoVec K k → BitReverseExecution K k
| 0, a => ⟨a, 1⟩
| _k + 1, a =>
let evenRun := bitReverseExec (evenCoeffs a)
let oddRun := bitReverseExec (oddCoeffs a)
⟨joinHalves evenRun.value oddRun.value,
evenRun.moves + oddRun.moves⟩Value projection of the bit-reversal execution.
def bitReverseCopy {K : Type*} {k : Nat} (a : PowTwoVec K k) :
PowTwoVec K k :=
(bitReverseExec a).valueA successor-size copy joins recursive copies of even and odd coefficients.
@[simp] theorem bitReverseCopy_succ {K : Type*} {k : Nat}
(a : PowTwoVec K (k + 1)) :
bitReverseCopy a =
joinHalves (bitReverseCopy (evenCoeffs a))
(bitReverseCopy (oddCoeffs a)) := rflThe copied value at a reversed index is the original value.
theorem bitReverseCopy_apply {K : Type*} {k : Nat}
(a : PowTwoVec K k) (i : Fin (2 ^ k)) :
bitReverseCopy a (bitReverseEquiv k i) = a i := by
induction k with
| zero =>
have hi : i = ⟨0, by norm_num⟩ := Fin.ext (by omega)
subst i
rfl
| succ k ih =>
obtain ⟨q, hq | hq⟩ := Nat.even_or_odd' i.1
· have hq_lt : q < 2 ^ k := by
have hi := i.2
simp [pow_succ] at hi
omega
let qi : Fin (2 ^ k) := ⟨q, hq_lt⟩
have hi : i = evenIndex qi := Fin.ext hq
subst i
rw [bitReverseEquiv_even, bitReverseCopy_succ, joinHalves_lower]
simpa using ih (evenCoeffs a) qi
· have hq_lt : q < 2 ^ k := by
have hi := i.2
simp [pow_succ] at hi
omega
let qi : Fin (2 ^ k) := ⟨q, hq_lt⟩
have hi : i = oddIndex qi := Fin.ext hq
subst i
rw [bitReverseEquiv_odd, bitReverseCopy_succ, joinHalves_upper]
simpa using ih (oddCoeffs a) qiApplying bit-reversal copy twice returns the original vector.
theorem bitReverseCopy_involutive {K : Type*} {k : Nat} :
Function.Involutive (@bitReverseCopy K k) := by
intro a
funext i
have hsem := bitReverseCopy_apply (bitReverseCopy a) (bitReverseEquiv k i)
rw [bitReverseEquiv_involutive k i] at hsem
exact hsem.trans (bitReverseCopy_apply a i)Bit-reversal performs exactly one functional output move per element.
@[simp] theorem bitReverseExec_moves {K : Type*} {k : Nat}
(a : PowTwoVec K k) :
(bitReverseExec a).moves = 2 ^ k := by
induction k with
| zero => rfl
| succ k ih =>
simp [bitReverseExec, ih, pow_succ]
omegaend Chapter30end CLRSCLRSLean.FourthEdition.Chapter_30.Section_30_3_Efficient_FFT_Implementations.IterativeFFT.Definitions
Chapter 30.3: Executable iterative radix-2 FFT
The complete transform folds globally ordered stages over a bit-reversed
vector. A stage recursively enumerates its independent contiguous blocks and
uses the already verified butterfly execution at each block. The arithmetic
fields charge the shared butterfly/twiddle schedule. Computing child roots
omega ^ 2 in nonfinal stages is root setup, excluded from these fields;
they do not count every ring operation in the literal function evaluator.
namespace CLRSnamespace Chapter30Extract the contiguous lower half of a successor-size vector.
def lowerHalf {K : Type*} {k : Nat} (a : PowTwoVec K (k + 1)) :
PowTwoVec K k := fun i => a (lowerHalfIndex i)Extract the contiguous upper half of a successor-size vector.
def upperHalf {K : Type*} {k : Nat} (a : PowTwoVec K (k + 1)) :
PowTwoVec K k := fun i => a (upperHalfIndex i)Reading a lower half uses the canonical lower-half embedding.
@[simp] theorem lowerHalf_apply {K : Type*} {k : Nat}
(a : PowTwoVec K (k + 1)) (i : Fin (2 ^ k)) :
lowerHalf a i = a (lowerHalfIndex i) := rflReading an upper half uses the canonical upper-half embedding.
@[simp] theorem upperHalf_apply {K : Type*} {k : Nat}
(a : PowTwoVec K (k + 1)) (i : Fin (2 ^ k)) :
upperHalf a i = a (upperHalfIndex i) := rflExtracting the lower half of joined vectors returns the lower input.
@[simp] theorem lowerHalf_joinHalves {K : Type*} {k : Nat}
(lower upper : PowTwoVec K k) :
lowerHalf (joinHalves lower upper) = lower := by
funext i
simp [lowerHalf]Extracting the upper half of joined vectors returns the upper input.
@[simp] theorem upperHalf_joinHalves {K : Type*} {k : Nat}
(lower upper : PowTwoVec K k) :
upperHalf (joinHalves lower upper) = upper := by
funext i
simp [upperHalf]Value and arithmetic counters for one global iterative FFT stage.
structure FFTStageExecution (K : Type*) (k : Nat) where
value : PowTwoVec K k
addSubtractions : Nat
multiplications : NatExecute one indexed stage. The final stage is one full butterfly layer; every earlier stage acts independently on the two halves with squared root.
def fftStageExec [Ring K] :
{k : Nat} → K → PowTwoVec K k → Fin k → FFTStageExecution K k
| 0, _, _, s => Fin.elim0 s
| k + 1, omega, a, s =>
if hfinal : s.1 = k then
let layer := butterflyLayerExec omega (lowerHalf a) (upperHalf a)
⟨layer.value, layer.addSubtractions, layer.multiplications⟩
else
let childStage : Fin k := ⟨s.1, by omega⟩
let lowerRun := fftStageExec (omega ^ 2) (lowerHalf a) childStage
let upperRun := fftStageExec (omega ^ 2) (upperHalf a) childStage
⟨joinHalves lowerRun.value upperRun.value,
lowerRun.addSubtractions + upperRun.addSubtractions,
lowerRun.multiplications + upperRun.multiplications⟩Value projection of one global stage.
def fftStage [Ring K] {k : Nat} (omega : K)
(a : PowTwoVec K k) (s : Fin k) : PowTwoVec K k :=
(fftStageExec omega a s).valueThe final indexed stage is one full butterfly layer.
@[simp] theorem fftStage_final [Ring K] {k : Nat} (omega : K)
(a : PowTwoVec K (k + 1)) :
fftStage omega a (Fin.last k) =
butterflyLayer omega (lowerHalf a) (upperHalf a) := by
simp [fftStage, fftStageExec, butterflyLayer]A nonfinal indexed stage acts independently on the two halves.
@[simp] theorem fftStage_nonfinal [Ring K] {k : Nat} (omega : K)
(a : PowTwoVec K (k + 1)) (s : Fin k) :
fftStage omega a s.castSucc =
joinHalves
(fftStage (omega ^ 2) (lowerHalf a) s)
(fftStage (omega ^ 2) (upperHalf a) s) := by
simp [fftStage, fftStageExec, Fin.castSucc, Nat.ne_of_lt s.2]Value and accumulated counters after an ordered stage prefix.
structure FFTStageSequenceExecution (K : Type*) (k : Nat) where
value : PowTwoVec K k
addSubtractions : Nat
multiplications : NatExecute the requested initial number of stages in increasing order.
def runFFTStagePrefixExec [Ring K] {k : Nat} (omega : K)
(a : PowTwoVec K k) :
(m : Nat) → m ≤ k → FFTStageSequenceExecution K k
| 0, _ => ⟨a, 0, 0⟩
| m + 1, hm =>
let previous := runFFTStagePrefixExec omega a m (by omega)
let current := fftStageExec omega previous.value ⟨m, by omega⟩
⟨current.value,
previous.addSubtractions + current.addSubtractions,
previous.multiplications + current.multiplications⟩Value projection after an ordered stage prefix.
def runFFTStagePrefix [Ring K] {k : Nat} (omega : K)
(a : PowTwoVec K k) (m : Nat) (hm : m ≤ k) : PowTwoVec K k :=
(runFFTStagePrefixExec omega a m hm).valueThe empty ordered stage prefix leaves its input unchanged.
@[simp] theorem runFFTStagePrefix_zero [Ring K] {k : Nat} (omega : K)
(a : PowTwoVec K k) (hm : 0 ≤ k) :
runFFTStagePrefix omega a 0 hm = a := rflExtending a stage prefix applies the next admissible stage.
theorem runFFTStagePrefix_succ [Ring K] {k m : Nat} (omega : K)
(a : PowTwoVec K k) (hm : m + 1 ≤ k) :
runFFTStagePrefix omega a (m + 1) hm =
fftStage omega
(runFFTStagePrefix omega a m (by omega)) ⟨m, by omega⟩ := rflExecute all admissible stages.
def runAllFFTStagesExec [Ring K] {k : Nat} (omega : K)
(a : PowTwoVec K k) : FFTStageSequenceExecution K k :=
runFFTStagePrefixExec omega a k le_rflValue projection after all stages.
def runAllFFTStages [Ring K] {k : Nat} (omega : K)
(a : PowTwoVec K k) : PowTwoVec K k :=
(runAllFFTStagesExec omega a).valueA complete successor-size run ends with the final stage.
theorem runAllFFTStages_succ [Ring K] {k : Nat} (omega : K)
(a : PowTwoVec K (k + 1)) :
runAllFFTStages omega a =
fftStage omega
(runFFTStagePrefix omega a k (Nat.le_succ k)) (Fin.last k) := by
exact runFFTStagePrefix_succ omega a le_rflResult and counters of bit reversal followed by all ordered stages.
structure IterativeFFTExecution (K : Type*) (k : Nat) where
value : PowTwoVec K k
bitReversalMoves : Nat
addSubtractions : Nat
multiplications : NatThe functional iterative radix-2 FFT execution.
def iterativeRadix2FFTExec [Ring K] {k : Nat} (omega : K)
(a : PowTwoVec K k) : IterativeFFTExecution K k :=
let reversal := bitReverseExec a
let stages := runAllFFTStagesExec omega reversal.value
⟨stages.value, reversal.moves,
stages.addSubtractions, stages.multiplications⟩Value projection of the iterative radix-2 execution.
def iterativeRadix2FFT [Ring K] {k : Nat} (omega : K)
(a : PowTwoVec K k) : PowTwoVec K k :=
(iterativeRadix2FFTExec omega a).valueThe iterative FFT is bit-reversal followed by all ordered stages.
theorem iterativeRadix2FFT_eq_runAll [Ring K] {k : Nat} (omega : K)
(a : PowTwoVec K k) :
iterativeRadix2FFT omega a =
runAllFFTStages omega (bitReverseCopy a) := rflErasing iterative execution counters yields the public transform value.
theorem iterativeRadix2FFTExec_value [Ring K] {k : Nat} (omega : K)
(a : PowTwoVec K k) :
(iterativeRadix2FFTExec omega a).value =
iterativeRadix2FFT omega a := rflThe singleton iterative transform is the identity.
@[simp] theorem iterativeRadix2FFT_zero [Ring K] (omega : K)
(a : PowTwoVec K 0) :
iterativeRadix2FFT omega a = a := by
funext i
have hi : i = ⟨0, by norm_num⟩ := Fin.ext (by omega)
subst i
rflend Chapter30end CLRSCLRSLean.FourthEdition.Chapter_30.Section_30_3_Efficient_FFT_Implementations.IterativeFFT.Correctness
Chapter 30.3: Iterative FFT correctness
The ordered stage prefix factors over the two contiguous halves until the final stage. That invariant gives the recursive FFT equation and transfers the proved recursive transform to the generic DFT.
namespace CLRSnamespace Chapter30A lifted child stage distributes over two joined halves.
theorem fftStage_join_castSucc [Ring K] {k : Nat} (omega : K)
(lower upper : PowTwoVec K k) (s : Fin k) :
fftStage omega (joinHalves lower upper) s.castSucc =
joinHalves (fftStage (omega ^ 2) lower s)
(fftStage (omega ^ 2) upper s) := by
rw [fftStage_nonfinal]
simpEvery requested nonfinal stage prefix acts independently on the two halves.
theorem runFFTStagePrefix_join [Ring K] {k m : Nat} (hm : m ≤ k)
(omega : K) (lower upper : PowTwoVec K k) :
runFFTStagePrefix omega (joinHalves lower upper) m
(hm.trans (Nat.le_succ k)) =
joinHalves
(runFFTStagePrefix (omega ^ 2) lower m hm)
(runFFTStagePrefix (omega ^ 2) upper m hm) := by
induction m with
| zero => rfl
| succ m ih =>
rw [runFFTStagePrefix_succ, runFFTStagePrefix_succ,
runFFTStagePrefix_succ]
rw [ih (by omega)]
exact fftStage_join_castSucc omega _ _ ⟨m, by omega⟩Specialization of the stage invariant to all child-size stages.
theorem runInitialFFTStages_join [Ring K] {k : Nat} (omega : K)
(lower upper : PowTwoVec K k) :
runFFTStagePrefix omega (joinHalves lower upper) k (Nat.le_succ k) =
joinHalves
(runAllFFTStages (omega ^ 2) lower)
(runAllFFTStages (omega ^ 2) upper) := by
simpa [runAllFFTStages, runAllFFTStagesExec, runFFTStagePrefix] using
runFFTStagePrefix_join (m := k) le_rfl omega lower upperThe iterative transform exposes the same even/odd recurrence as the recursive radix-2 FFT.
theorem iterativeRadix2FFT_succ [Ring K] {k : Nat} (omega : K)
(a : PowTwoVec K (k + 1)) :
iterativeRadix2FFT omega a =
butterflyLayer omega
(iterativeRadix2FFT (omega ^ 2) (evenCoeffs a))
(iterativeRadix2FFT (omega ^ 2) (oddCoeffs a)) := by
rw [iterativeRadix2FFT_eq_runAll, runAllFFTStages_succ,
bitReverseCopy_succ, runInitialFFTStages_join, fftStage_final]
rw [iterativeRadix2FFT_eq_runAll, iterativeRadix2FFT_eq_runAll]
simpThe ordered iterative execution has the same value as the canonical recursive radix-2 execution for every ring and every supplied root.
theorem iterativeRadix2FFT_eq_recursiveFFT [Ring K] {k : Nat}
(omega : K) (a : PowTwoVec K k) :
iterativeRadix2FFT omega a = recursiveFFT omega a := by
induction k generalizing omega with
| zero =>
funext i
have hi : i = ⟨0, by norm_num⟩ := Fin.ext (by omega)
subst i
rfl
| succ k ih =>
rw [iterativeRadix2FFT_succ]
by_cases hk : k = 0
· subst k
simp [recursiveFFT, recursiveFFTExec, butterflyLayer,
butterflyLayerExec]
· have hchildRoot :
twiddleChildRoot k omega
(twiddlePowersAuxExec omega (2 ^ k) 1) = omega ^ 2 :=
twiddleChildRoot_eq_square (Nat.pos_of_ne_zero hk) omega
simp only [recursiveFFT, recursiveFFTExec]
rw [hchildRoot]
change butterflyLayer omega
(iterativeRadix2FFT (omega ^ 2) (evenCoeffs a))
(iterativeRadix2FFT (omega ^ 2) (oddCoeffs a)) =
butterflyLayer omega
(recursiveFFT (omega ^ 2) (evenCoeffs a))
(recursiveFFT (omega ^ 2) (oddCoeffs a))
rw [ih, ih]Under the existing primitive-root hypotheses, the iterative FFT computes the generic DFT.
theorem iterativeRadix2FFT_eq_dft [Field K] [CharZero K] {k : Nat}
{omega : K} (homega : IsPrimitiveRoot omega (2 ^ k))
(a : PowTwoVec K k) :
iterativeRadix2FFT omega a = dft omega a := by
rw [iterativeRadix2FFT_eq_recursiveFFT]
exact recursiveFFT_eq_dft homega aend Chapter30end CLRSCLRSLean.FourthEdition.Chapter_30.Section_30_3_Efficient_FFT_Implementations.IterativeFFT.Costs
Chapter 30.3: Iterative FFT costs
All counts below are projections of the value-producing bit-reversal and stage executions in their stated charge model. Data movement is separate from field arithmetic. Nonfinal-stage child-root squaring is excluded setup work. Each butterfly counts a product once under shared arithmetic even though the function-valued output mentions it in both sum and difference. These are not counts of every operation in the literal Lean evaluator.
namespace CLRSnamespace Chapter30Total charged arithmetic work of one iterative stage execution.
def FFTStageExecution.work (r : FFTStageExecution K k) : Nat :=
r.addSubtractions + r.multiplicationsCharged additions/subtractions and multiplications of an iterative FFT.
def IterativeFFTExecution.arithmeticWork
(r : IterativeFFTExecution K k) : Nat :=
r.addSubtractions + r.multiplicationsCharged iterative work including bit-reversal movement, excluding child-root setup and assuming shared butterfly products.
def IterativeFFTExecution.totalWork
(r : IterativeFFTExecution K k) : Nat :=
r.bitReversalMoves + r.arithmeticWorkEvery global stage charges one addition/subtraction per output slot.
@[simp] theorem fftStageExec_addSubtractions [Ring K] {k : Nat}
(omega : K) (a : PowTwoVec K k) (s : Fin k) :
(fftStageExec omega a s).addSubtractions = 2 ^ k := by
induction k generalizing omega with
| zero => exact Fin.elim0 s
| succ k ih =>
by_cases h : s.1 = k
· simp [fftStageExec, h, butterflyLayerExec_addSubtractions, pow_succ]
omega
· have hs : s.1 < k := by omega
simp [fftStageExec, h, ih, pow_succ]
omegaEvery global stage charges one multiplication per output slot.
@[simp] theorem fftStageExec_multiplications [Ring K] {k : Nat}
(omega : K) (a : PowTwoVec K k) (s : Fin k) :
(fftStageExec omega a s).multiplications = 2 ^ k := by
induction k generalizing omega with
| zero => exact Fin.elim0 s
| succ k ih =>
by_cases h : s.1 = k
· simp [fftStageExec, h, butterflyLayerExec_multiplications, pow_succ]
omega
· have hs : s.1 < k := by omega
simp [fftStageExec, h, ih, pow_succ]
omega
An m-stage prefix charges m * 2^k additions/subtractions.
@[simp] theorem runFFTStagePrefixExec_addSubtractions [Ring K]
{k m : Nat} (omega : K) (a : PowTwoVec K k) (hm : m ≤ k) :
(runFFTStagePrefixExec omega a m hm).addSubtractions = m * 2 ^ k := by
induction m with
| zero => simp [runFFTStagePrefixExec]
| succ m ih =>
simp [runFFTStagePrefixExec, ih, Nat.succ_mul]
An m-stage prefix charges m * 2^k multiplications.
@[simp] theorem runFFTStagePrefixExec_multiplications [Ring K]
{k m : Nat} (omega : K) (a : PowTwoVec K k) (hm : m ≤ k) :
(runFFTStagePrefixExec omega a m hm).multiplications = m * 2 ^ k := by
induction m with
| zero => simp [runFFTStagePrefixExec]
| succ m ih =>
simp [runFFTStagePrefixExec, ih, Nat.succ_mul]A complete iterative execution charges one bit-reversal move per input.
@[simp] theorem iterativeRadix2FFTExec_bitReversalMoves [Ring K] {k : Nat}
(omega : K) (a : PowTwoVec K k) :
(iterativeRadix2FFTExec omega a).bitReversalMoves = 2 ^ k := by
simp [iterativeRadix2FFTExec]
A complete iterative execution charges k * 2^k additions/subtractions.
@[simp] theorem iterativeRadix2FFTExec_addSubtractions [Ring K] {k : Nat}
(omega : K) (a : PowTwoVec K k) :
(iterativeRadix2FFTExec omega a).addSubtractions = k * 2 ^ k := by
simp [iterativeRadix2FFTExec, runAllFFTStagesExec]
A complete iterative execution charges k * 2^k multiplications.
@[simp] theorem iterativeRadix2FFTExec_multiplications [Ring K] {k : Nat}
(omega : K) (a : PowTwoVec K k) :
(iterativeRadix2FFTExec omega a).multiplications = k * 2 ^ k := by
simp [iterativeRadix2FFTExec, runAllFFTStagesExec]Iterative arithmetic work equals the recursive radix-2 cost function.
theorem iterativeRadix2FFTExec_arithmeticWork [Ring K] {k : Nat}
(omega : K) (a : PowTwoVec K k) :
(iterativeRadix2FFTExec omega a).arithmeticWork = radix2FFTWork k := by
simp [IterativeFFTExecution.arithmeticWork, radix2FFTWork, Nat.two_mul]
rw [Nat.add_mul]Total iterative work, including one bit-reversal move per element.
def iterativeRadix2FFTTotalWork (k : Nat) : Nat :=
2 ^ k + radix2FFTWork kThe recursive arithmetic cost exposes its exact closed form.
@[simp] theorem radix2FFTWork_closed (k : Nat) :
radix2FFTWork k = 2 * k * 2 ^ k := rflIterative total work is bit-reversal movement plus exact arithmetic work.
@[simp] theorem iterativeRadix2FFTTotalWork_closed (k : Nat) :
iterativeRadix2FFTTotalWork k = 2 ^ k + 2 * k * 2 ^ k := rflThe execution's total-work field equals the public iterative cost.
theorem iterativeRadix2FFTExec_totalWork [Ring K] {k : Nat}
(omega : K) (a : PowTwoVec K k) :
(iterativeRadix2FFTExec omega a).totalWork =
iterativeRadix2FFTTotalWork k := by
simp [IterativeFFTExecution.totalWork, iterativeRadix2FFTTotalWork,
iterativeRadix2FFTExec_arithmeticWork]Bit-reversal moves do not change the exact-power FFT asymptotic scale.
theorem iterativeRadix2FFTTotalWork_bigTheta :
Chapter03.isBigTheta
(fun k => (iterativeRadix2FFTTotalWork k : ℝ))
(fun k => (k : ℝ) * (2 : ℝ) ^ k) := by
constructor
· refine (Chapter03.isBigO_iff _ _).mpr ⟨3, by norm_num, 1, ?_⟩
intro k hk
rw [abs_of_nonneg (Nat.cast_nonneg _), abs_of_nonneg (by positivity)]
simp [iterativeRadix2FFTTotalWork, radix2FFTWork]
have hk' : (1 : ℝ) ≤ k := by exact_mod_cast hk
have hpow : 0 ≤ (2 : ℝ) ^ k := by positivity
nlinarith
· refine (Chapter03.isBigOmega_iff _ _).mpr ⟨2, by norm_num, 0, ?_⟩
intro k _
rw [abs_of_nonneg (by positivity), abs_of_nonneg (Nat.cast_nonneg _)]
simp [iterativeRadix2FFTTotalWork, radix2FFTWork]
have hpow : 0 ≤ (2 : ℝ) ^ k := by positivity
nlinarithTotal iterative work at the least nonempty padded FFT capacity.
def paddedIterativeFFTWork (n : Nat) : Nat :=
fftCapacity n + paddedFFTWork nPadded iterative work is monotone in the advertised input capacity.
theorem paddedIterativeFFTWork_monotone : Monotone paddedIterativeFFTWork := by
intro m n hmn
exact Nat.add_le_add (fftCapacity_monotone hmn) (paddedFFTWork_monotone hmn)Padded total work remains attached to the actual iterative execution.
theorem iterativeRadix2FFTExec_zeroPad_totalWork [Ring K] {n : Nat}
(omega : K) (a : CoeffVector K n) :
(iterativeRadix2FFTExec (k := fftExponent n) omega
(zeroPadToFFTCapacity a)).totalWork = paddedIterativeFFTWork n := by
rw [iterativeRadix2FFTExec_totalWork]
rflInput capacity at least two selects a positive FFT exponent.
private theorem fftExponent_pos_of_two_le {n : Nat} (hn : 2 ≤ n) :
0 < fftExponent n := by
rw [fftExponent, max_eq_right (by omega)]
exact Nat.clog_pos (by norm_num) (by omega)Above singleton inputs, arithmetic work dominates padded capacity.
private theorem fftCapacity_le_paddedFFTWork {n : Nat} (hn : 2 ≤ n) :
fftCapacity n ≤ paddedFFTWork n := by
have hexp : 1 ≤ fftExponent n := fftExponent_pos_of_two_le hn
have hfactor : 1 ≤ 2 * fftExponent n := by omega
unfold fftCapacity paddedFFTWork radix2FFTWork
simpa [Nat.mul_assoc] using
Nat.mul_le_mul_right (2 ^ fftExponent n) hfactorPadded recursive arithmetic work is bounded by iterative total work.
theorem paddedFFTWork_le_paddedIterativeFFTWork (n : Nat) :
paddedFFTWork n ≤ paddedIterativeFFTWork n := by
simp [paddedIterativeFFTWork]Above singleton inputs, iterative total work is at most twice arithmetic work.
theorem paddedIterativeFFTWork_le_two_mul {n : Nat} (hn : 2 ≤ n) :
paddedIterativeFFTWork n ≤ 2 * paddedFFTWork n := by
have hcap := fftCapacity_le_paddedFFTWork hn
unfold paddedIterativeFFTWork
omegaPadded iterative total work has the same asymptotic scale as arithmetic work.
private theorem paddedIterativeFFTWork_isBigTheta_paddedFFTWork :
Chapter03.isBigTheta
(fun n : Nat => (paddedIterativeFFTWork n : ℝ))
(fun n : Nat => (paddedFFTWork n : ℝ)) := by
constructor
· refine (Chapter03.isBigO_iff _ _).mpr ⟨2, by norm_num, 2, ?_⟩
intro n hn
rw [abs_of_nonneg (Nat.cast_nonneg _), abs_of_nonneg (Nat.cast_nonneg _)]
exact_mod_cast paddedIterativeFFTWork_le_two_mul hn
· refine (Chapter03.isBigOmega_iff _ _).mpr ⟨1, by norm_num, 0, ?_⟩
intro n _
rw [abs_of_nonneg (Nat.cast_nonneg _), abs_of_nonneg (Nat.cast_nonneg _)]
norm_num
exact_mod_cast paddedFFTWork_le_paddedIterativeFFTWork nPadding lifts the iterative execution, including bit-reversal moves, to the all-input textbook linear-logarithmic scale.
theorem paddedIterativeFFTWork_allInput_bigTheta :
Chapter03.isBigTheta
(fun n : Nat => (paddedIterativeFFTWork n : ℝ))
(Chapter04.realLogLogScale 2 2) := by
exact Chapter03.isBigTheta_trans
paddedIterativeFFTWork_isBigTheta_paddedFFTWork
paddedFFTWork_allInput_bigThetaend Chapter30end CLRSCLRSLean.FourthEdition.Chapter_30.Section_30_3_Efficient_FFT_Implementations.ParallelFFT
Chapter 30.3: Layered parallel FFT network
The network exposes one typed layer per iterative stage. Each layer stores a recursive stage circuit whose leaves are actual butterfly gates with fixed twiddle constants; evaluation, butterfly count, and depth all interpret that same syntax. Bit-reversal is wiring rather than an arithmetic gate.
namespace CLRSnamespace Chapter30One logical radix-2 butterfly with a fixed twiddle constant.
structure FFTButterflyGate (K : Type*) where
twiddle : KEvaluate one logical butterfly on its lower and upper inputs.
def FFTButterflyGate.eval [Ring K] (gate : FFTButterflyGate K)
(u v : K) : K × K :=
let product := gate.twiddle * v
(u + product, u - product)A complete local butterfly layer, with one actual gate at every offset.
structure ButterflyLayerCircuit (K : Type*) (k : Nat) where
gates : Fin (2 ^ k) → FFTButterflyGate KEvaluate every stored gate in a local butterfly layer in parallel.
def ButterflyLayerCircuit.eval [Ring K] (layer : ButterflyLayerCircuit K k)
(u v : PowTwoVec K k) : PowTwoVec K (k + 1) :=
joinHalves
(fun j => (layer.gates j).eval (u j) (v j) |>.1)
(fun j => (layer.gates j).eval (u j) (v j) |>.2)
The canonical local layer stores twiddle omega ^ j at gate j.
def canonicalButterflyLayerCircuit [Monoid K] (omega : K) (k : Nat) :
ButterflyLayerCircuit K k :=
⟨fun j => ⟨omega ^ j.1⟩⟩Joining the two extracted halves reconstructs the original vector.
private theorem joinHalves_lowerHalf_upperHalf {K : Type*} {k : Nat}
(a : PowTwoVec K (k + 1)) :
joinHalves (lowerHalf a) (upperHalf a) = a := by
funext i
have h : ∀ t : Fin (2 ^ k + 2 ^ k),
joinHalves (lowerHalf a) (upperHalf a) ((powTwoSuccEquiv k).symm t) =
a ((powTwoSuccEquiv k).symm t) := by
intro t
refine Fin.addCases ?_ ?_ t
· intro j
exact joinHalves_lower (lowerHalf a) (upperHalf a) j
· intro j
exact joinHalves_upper (lowerHalf a) (upperHalf a) j
simpa using h (powTwoSuccEquiv k i)Equality of both contiguous halves determines a successor-size vector.
private theorem powTwoVec_eq_of_halves {K : Type*} {k : Nat}
{a b : PowTwoVec K (k + 1)}
(hlower : lowerHalf a = lowerHalf b)
(hupper : upperHalf a = upperHalf b) : a = b := by
rw [← joinHalves_lowerHalf_upperHalf a,
← joinHalves_lowerHalf_upperHalf b, hlower, hupper]Evaluating the canonical stored gates is exactly the verified butterfly layer from the recursive FFT.
theorem canonicalButterflyLayerCircuit_eval [Ring K] {k : Nat} (omega : K)
(u v : PowTwoVec K k) :
(canonicalButterflyLayerCircuit omega k).eval u v =
butterflyLayer omega u v := by
apply powTwoVec_eq_of_halves
· funext j
simp [ButterflyLayerCircuit.eval, canonicalButterflyLayerCircuit,
FFTButterflyGate.eval]
· funext j
simp [ButterflyLayerCircuit.eval, canonicalButterflyLayerCircuit,
FFTButterflyGate.eval]Syntax for one global FFT stage: either one complete local butterfly layer, or two equal-depth stage circuits evaluated independently on the two halves.
inductive FFTStageCircuit (K : Type*) : Nat → Type _
| butterfly {k : Nat} (layer : ButterflyLayerCircuit K k) :
FFTStageCircuit K (k + 1)
| parallel {k : Nat} (lower upper : FFTStageCircuit K k) :
FFTStageCircuit K (k + 1)Evaluate the stored syntax of one global FFT stage.
def FFTStageCircuit.eval [Ring K] :
{k : Nat} → FFTStageCircuit K k → PowTwoVec K k → PowTwoVec K k
| _, .butterfly layer, a => layer.eval (lowerHalf a) (upperHalf a)
| _, .parallel lower upper, a =>
joinHalves (lower.eval (lowerHalf a)) (upper.eval (upperHalf a))Number of actual logical butterflies stored in a local gate family.
def ButterflyLayerCircuit.butterflyCount
(_layer : ButterflyLayerCircuit K k) : Nat :=
Fintype.card (Fin (2 ^ k))Structural butterfly count of one stage circuit.
def FFTStageCircuit.butterflyCount : {k : Nat} → FFTStageCircuit K k → Nat
| _, .butterfly layer => layer.butterflyCount
| _, .parallel lower upper => lower.butterflyCount + upper.butterflyCountStructural butterfly depth of one stage circuit.
def FFTStageCircuit.butterflyDepth : {k : Nat} → FFTStageCircuit K k → Nat
| _, .butterfly _ => 1
| _, .parallel lower upper => max lower.butterflyDepth upper.butterflyDepthConstruct the canonical stored circuit for one globally indexed stage.
def fftStageCircuit [Monoid K] :
{k : Nat} → K → Fin k → FFTStageCircuit K k
| 0, _, s => Fin.elim0 s
| k + 1, omega, s =>
if hfinal : s.1 = k then
.butterfly (canonicalButterflyLayerCircuit omega k)
else
let childStage : Fin k := ⟨s.1, by omega⟩
.parallel (fftStageCircuit (omega ^ 2) childStage)
(fftStageCircuit (omega ^ 2) childStage)The canonical stored stage circuit evaluates to the verified iterative stage semantics.
theorem fftStageCircuit_eval [Ring K] {k : Nat} (omega : K)
(a : PowTwoVec K k) (s : Fin k) :
(fftStageCircuit omega s).eval a = fftStage omega a s := by
induction k generalizing omega with
| zero => exact Fin.elim0 s
| succ k ih =>
by_cases hfinal : s.1 = k
· have hs : s = Fin.last k := Fin.ext hfinal
subst s
simp [fftStageCircuit, FFTStageCircuit.eval,
canonicalButterflyLayerCircuit_eval]
· let childStage : Fin k := Fin.castLT s (by omega)
have hs : childStage.castSucc = s := Fin.ext rfl
rw [← hs]
simp [fftStageCircuit, FFTStageCircuit.eval, ih,
Nat.ne_of_lt childStage.2]
Every canonical global stage stores exactly 2^(k-1) butterflies.
theorem fftStageCircuit_butterflyCount {K : Type*} [Monoid K] {k : Nat}
(omega : K) (s : Fin k) :
(fftStageCircuit omega s).butterflyCount = 2 ^ (k - 1) := by
induction k generalizing omega with
| zero => exact Fin.elim0 s
| succ k ih =>
by_cases hfinal : s.1 = k
· simp [fftStageCircuit, hfinal, FFTStageCircuit.butterflyCount,
ButterflyLayerCircuit.butterflyCount]
· have hk : 0 < k := by omega
let childStage : Fin k := Fin.castLT s (by omega)
have hs : childStage.castSucc = s := Fin.ext rfl
rw [← hs]
simp [fftStageCircuit, FFTStageCircuit.butterflyCount, ih,
Nat.ne_of_lt childStage.2]
have hk1 : 1 ≤ k := hk
calc
2 ^ (k - 1) + 2 ^ (k - 1) = 2 ^ (k - 1) * 2 := by omega
_ = 2 ^ ((k - 1) + 1) := by rw [pow_succ]
_ = 2 ^ k := by rw [Nat.sub_add_cancel hk1]Every canonical global stage has one butterfly layer of structural depth.
theorem fftStageCircuit_butterflyDepth {K : Type*} [Monoid K] {k : Nat}
(omega : K) (s : Fin k) :
(fftStageCircuit omega s).butterflyDepth = 1 := by
induction k generalizing omega with
| zero => exact Fin.elim0 s
| succ k ih =>
by_cases hfinal : s.1 = k
· simp [fftStageCircuit, hfinal, FFTStageCircuit.butterflyDepth]
· let childStage : Fin k := Fin.castLT s (by omega)
have hs : childStage.castSucc = s := Fin.ext rfl
rw [← hs]
simp [fftStageCircuit, FFTStageCircuit.butterflyDepth, ih,
Nat.ne_of_lt childStage.2]A butterfly is identified by its contiguous block and its within-half offset at one stage.
abbrev FFTButterflyPosition (k : Nat) (s : Fin k) :=
Fin (2 ^ (k - s.1 - 1)) × Fin (2 ^ s.1)One logical butterfly layer of the canonical network.
structure FFTLayer (K : Type*) (k : Nat) where
omega : K
stage : Fin k
circuit : FFTStageCircuit K kConstruct the canonical stored layer for one stage.
def fftLayer [Monoid K] (omega : K) {k : Nat} (s : Fin k) : FFTLayer K k :=
⟨omega, s, fftStageCircuit omega s⟩Evaluate the actual stage circuit stored by a layer.
def FFTLayer.eval [Ring K] (layer : FFTLayer K k)
(a : PowTwoVec K k) : PowTwoVec K k :=
layer.circuit.eval aThe stage root used by every block in this layer.
def FFTLayer.root [Monoid K] (layer : FFTLayer K k) : K :=
layer.omega ^ (2 ^ (k - layer.stage.1 - 1))The fixed twiddle constant at one butterfly position.
def FFTLayer.twiddle [Monoid K] (layer : FFTLayer K k)
(position : FFTButterflyPosition k layer.stage) : K :=
layer.root ^ position.2.1A typed family of logical FFT layers.
structure FFTNetwork (K : Type*) (k : Nat) where
layers : Fin k → FFTLayer K kThe canonical network contains the ordered stages for one supplied root.
def fftNetwork {K : Type*} [Monoid K] {k : Nat} (omega : K) : FFTNetwork K k :=
⟨fun s => fftLayer omega s⟩Evaluate the requested prefix of a typed network.
def FFTNetwork.evalPrefix [Ring K] (network : FFTNetwork K k)
(a : PowTwoVec K k) : (m : Nat) → m ≤ k → PowTwoVec K k
| 0, _ => a
| m + 1, hm =>
let previous := network.evalPrefix a m (by omega)
let layer := network.layers ⟨m, by omega⟩
layer.eval previousEvaluate all arithmetic layers, without bit-reversal wiring.
def FFTNetwork.evalLayers [Ring K] (network : FFTNetwork K k)
(a : PowTwoVec K k) : PowTwoVec K k :=
network.evalPrefix a k le_rflEvaluate bit-reversal wiring followed by all arithmetic layers.
def FFTNetwork.eval [Ring K] (network : FFTNetwork K k)
(a : PowTwoVec K k) : PowTwoVec K k :=
network.evalLayers (bitReverseCopy a)Every canonical circuit prefix agrees with the functional stage prefix.
private theorem fftNetwork_evalPrefix [Ring K] {k m : Nat} (omega : K)
(a : PowTwoVec K k) (hm : m ≤ k) :
(fftNetwork omega).evalPrefix a m hm =
runFFTStagePrefix omega a m hm := by
induction m with
| zero => rfl
| succ m ih =>
change (fftStageCircuit omega ⟨m, by omega⟩).eval
((fftNetwork omega).evalPrefix a m (by omega)) =
fftStage omega
(runFFTStagePrefix omega a m (by omega)) ⟨m, by omega⟩
rw [fftStageCircuit_eval, ih]Evaluating all canonical arithmetic layers agrees with all FFT stages.
theorem fftNetwork_evalLayers [Ring K] {k : Nat} (omega : K)
(a : PowTwoVec K k) :
(fftNetwork omega).evalLayers a = runAllFFTStages omega a := by
exact fftNetwork_evalPrefix omega a le_rflThe explicit canonical network evaluates to the iterative FFT.
theorem fftNetwork_eval [Ring K] {k : Nat} (omega : K)
(a : PowTwoVec K k) :
(fftNetwork omega).eval a = iterativeRadix2FFT omega a := by
rw [FFTNetwork.eval, fftNetwork_evalLayers,
iterativeRadix2FFT_eq_runAll]Number of independent butterflies represented by one layer.
def FFTLayer.butterflyCount (layer : FFTLayer K k) : Nat :=
layer.circuit.butterflyCountSum of butterfly positions across all layers.
def FFTNetwork.butterflyCount (network : FFTNetwork K k) : Nat :=
∑ s : Fin k, (network.layers s).butterflyCountEvery canonical layer stores the exact per-stage butterfly count.
theorem fftLayer_butterflyCount {K : Type*} [Monoid K] {k : Nat}
(omega : K) (s : Fin k) :
(fftLayer omega s).butterflyCount = 2 ^ (k - 1) := by
exact fftStageCircuit_butterflyCount omega s
A length-2^k FFT network has k * 2^(k-1) butterflies.
theorem fftNetwork_butterflyCount {K : Type*} [Monoid K] {k : Nat} (omega : K) :
(fftNetwork omega : FFTNetwork K k).butterflyCount =
k * 2 ^ (k - 1) := by
change (∑ s : Fin k, (fftLayer omega s).butterflyCount) =
k * 2 ^ (k - 1)
simp [fftLayer_butterflyCount]Logical depth of one stored stage circuit.
def FFTLayer.butterflyDepth (layer : FFTLayer K k) : Nat :=
layer.circuit.butterflyDepthLogical depth obtained by composing the stored layer circuits.
def FFTNetwork.butterflyDepth (network : FFTNetwork K k) : Nat :=
∑ s : Fin k, (network.layers s).butterflyDepthOne multiplication and two addition/subtraction gates per butterfly.
def FFTNetwork.primitiveGateCount (network : FFTNetwork K k) : Nat :=
3 * network.butterflyCountEach butterfly layer expands to a multiplication level followed by an addition/subtraction level.
def FFTNetwork.primitiveDepth (network : FFTNetwork K k) : Nat :=
2 * network.butterflyDepth
A canonical k-layer network has butterfly depth exactly k.
@[simp] theorem fftNetwork_butterflyDepth {K : Type*} [Monoid K]
{k : Nat} (omega : K) :
(fftNetwork (k := k) omega).butterflyDepth = k := by
change (∑ s : Fin k, (fftLayer omega s).butterflyDepth) = k
simp [FFTLayer.butterflyDepth, fftLayer,
fftStageCircuit_butterflyDepth]Expanding canonical butterflies gives the exact primitive-gate count.
theorem fftNetwork_primitiveGateCount {K : Type*} [Monoid K]
{k : Nat} (omega : K) :
(fftNetwork (k := k) omega).primitiveGateCount =
3 * k * 2 ^ (k - 1) := by
rw [FFTNetwork.primitiveGateCount, fftNetwork_butterflyCount]
simp [Nat.mul_assoc]
Expanding each butterfly to two primitive levels gives depth 2 * k.
@[simp] theorem fftNetwork_primitiveDepth {K : Type*} [Monoid K]
{k : Nat} (omega : K) :
(fftNetwork (k := k) omega).primitiveDepth = 2 * k := by
simp [FFTNetwork.primitiveDepth]end Chapter30end CLRSScope and implementation notes
Imports
Current source
Sections 30.1--30.3 are native fourth-edition sections (representing
polynomials, the DFT and FFT, and efficient FFT implementations), imported
directly from
Section 30.1,
Section 30.2,
and
Section 30.3.
Declarations keep their current namespaces; the third-edition-numbered
imports CLRSLean.Chapter_30 and
CLRSLean.Chapter_30.Section_30_* forward to these sources.
Implementation details
The supporting implementation pages remain available outside the main sidebar:
Coverage boundary
The native sections supply the represented fourth-edition polynomial/FFT
sections (the FFT correctness and work analysis, the bit-reversal and
iterative-FFT implementations, and the parallel FFT). Exact arithmetic fields
use shared butterfly products; function-valued outputs do not establish
memoization when their two slots are evaluated separately. Iterative child-root
squaring is excluded setup work. Bit-reversal movement has a separate counter.
The parallel arithmetic-circuit model treats roots as constants. DFT correctness
and the stated asymptotic theorems hold in these respective charge models;
totalWork is not a count of every operation of the literal evaluator.
See docs/clrs-fourth-edition-map.csv for the section-level mapping and
docs/migrations/clrs4.md for compatibility and deprecation policy.
CLRS, fourth edition · Chapter 30 of 35