Imports
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 CLRS