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