Helper lemmas for ntt_outerLoop_computes_ref_ntt #
All entries of ensureRoots m are strictly below mod32.
ℕ-variant of ensure_roots_bound: all entries of ensureRoots m are < mod32.
ℕ-variant: if m = 2^n and n is even, then nttInplace.go 64 m 0 &&& 1 = 0.
ℕ-variant: if m = 2^n and n is odd, then nttInplace.go 64 m 0 &&& 1 ≠ 0.
theorem
inv_at_n_implies_ref_ntt
{m : ℕ}
(n : ℕ)
(hm_eq : m = 2 ^ n)
(v a : Vector UInt32 m)
(hinv : outerLoop_inv n n ⋯ hm_eq v a)
:
@[simp]
theorem
radix4Middle_zero_blocks
{n : ℕ}
(inverse : Bool)
(roots : Vector UInt32 n)
(s len b : ℕ)
(a : Vector UInt32 n)
:
When the block count is 0 (which happens when len is large enough relative to m
or when len = 0), radix4Middle does nothing.
bitRev radix-4 decomposition #
ntt_sub_input stride-4 relations #
ref_ntt radix-4 unfolding #
theorem
ref_ntt_radix4_q0
{R : Type u_1}
[CommRing R]
(q : ℕ)
(ω : R)
(f : Fin (2 ^ (q + 2)) → R)
(j2 : ℕ)
(hj2 : j2 < 2 ^ q)
:
ref_ntt (q + 2) ω f ⟨j2, ⋯⟩ = ref_ntt q (ω ^ 4) (fun (j : Fin (2 ^ q)) => f ⟨4 * ↑j, ⋯⟩) ⟨j2, hj2⟩ + (ω ^ 2) ^ j2 * ref_ntt q (ω ^ 4) (fun (j : Fin (2 ^ q)) => f ⟨4 * ↑j + 2, ⋯⟩) ⟨j2, hj2⟩ + ω ^ j2 * (ref_ntt q (ω ^ 4) (fun (j : Fin (2 ^ q)) => f ⟨4 * ↑j + 1, ⋯⟩) ⟨j2, hj2⟩ + (ω ^ 2) ^ j2 * ref_ntt q (ω ^ 4) (fun (j : Fin (2 ^ q)) => f ⟨4 * ↑j + 3, ⋯⟩) ⟨j2, hj2⟩)
theorem
ref_ntt_radix4_q1
{R : Type u_1}
[CommRing R]
(q : ℕ)
(ω : R)
(f : Fin (2 ^ (q + 2)) → R)
(j2 : ℕ)
(hj2 : j2 < 2 ^ q)
:
ref_ntt (q + 2) ω f ⟨j2 + 2 ^ q, ⋯⟩ = ref_ntt q (ω ^ 4) (fun (j : Fin (2 ^ q)) => f ⟨4 * ↑j, ⋯⟩) ⟨j2, hj2⟩ - (ω ^ 2) ^ j2 * ref_ntt q (ω ^ 4) (fun (j : Fin (2 ^ q)) => f ⟨4 * ↑j + 2, ⋯⟩) ⟨j2, hj2⟩ + ω ^ (j2 + 2 ^ q) * (ref_ntt q (ω ^ 4) (fun (j : Fin (2 ^ q)) => f ⟨4 * ↑j + 1, ⋯⟩) ⟨j2, hj2⟩ - (ω ^ 2) ^ j2 * ref_ntt q (ω ^ 4) (fun (j : Fin (2 ^ q)) => f ⟨4 * ↑j + 3, ⋯⟩) ⟨j2, hj2⟩)
theorem
ref_ntt_radix4_q2
{R : Type u_1}
[CommRing R]
(q : ℕ)
(ω : R)
(f : Fin (2 ^ (q + 2)) → R)
(j2 : ℕ)
(hj2 : j2 < 2 ^ q)
:
ref_ntt (q + 2) ω f ⟨j2 + 2 ^ (q + 1), ⋯⟩ = ref_ntt q (ω ^ 4) (fun (j : Fin (2 ^ q)) => f ⟨4 * ↑j, ⋯⟩) ⟨j2, hj2⟩ + (ω ^ 2) ^ j2 * ref_ntt q (ω ^ 4) (fun (j : Fin (2 ^ q)) => f ⟨4 * ↑j + 2, ⋯⟩) ⟨j2, hj2⟩ - ω ^ j2 * (ref_ntt q (ω ^ 4) (fun (j : Fin (2 ^ q)) => f ⟨4 * ↑j + 1, ⋯⟩) ⟨j2, hj2⟩ + (ω ^ 2) ^ j2 * ref_ntt q (ω ^ 4) (fun (j : Fin (2 ^ q)) => f ⟨4 * ↑j + 3, ⋯⟩) ⟨j2, hj2⟩)
Quadrant 2: position j2 + 2^(q+1).
theorem
ref_ntt_radix4_q3
{R : Type u_1}
[CommRing R]
(q : ℕ)
(ω : R)
(f : Fin (2 ^ (q + 2)) → R)
(j2 : ℕ)
(hj2 : j2 < 2 ^ q)
:
ref_ntt (q + 2) ω f ⟨j2 + 2 ^ q + 2 ^ (q + 1), ⋯⟩ = ref_ntt q (ω ^ 4) (fun (j : Fin (2 ^ q)) => f ⟨4 * ↑j, ⋯⟩) ⟨j2, hj2⟩ - (ω ^ 2) ^ j2 * ref_ntt q (ω ^ 4) (fun (j : Fin (2 ^ q)) => f ⟨4 * ↑j + 2, ⋯⟩) ⟨j2, hj2⟩ - ω ^ (j2 + 2 ^ q) * (ref_ntt q (ω ^ 4) (fun (j : Fin (2 ^ q)) => f ⟨4 * ↑j + 1, ⋯⟩) ⟨j2, hj2⟩ - (ω ^ 2) ^ j2 * ref_ntt q (ω ^ 4) (fun (j : Fin (2 ^ q)) => f ⟨4 * ↑j + 3, ⋯⟩) ⟨j2, hj2⟩)
theorem
ntt_block_pos_arith_nat
{m : ℕ}
(n q : ℕ)
(hq2 : q + 2 ≤ n)
(hm_eq : m = 2 ^ n)
(hn64 : n < 64)
(len : UInt64)
(hlen : len.toNat = 2 ^ (q + 1))
(b' j2' : ℕ)
(hb' : b' < 2 ^ (n - q - 2))
(hj2' : j2' < 2 ^ q)
:
have s := len >>> 1;
have i2' := (b' * 2 * len.toNat).toUInt64;
i2'.toNat = b' * 2 ^ (q + 2) ∧ (i2' + j2'.toUInt64).toNat = b' * 2 ^ (q + 2) + j2' ∧ (i2' + j2'.toUInt64 + s).toNat = b' * 2 ^ (q + 2) + j2' + 2 ^ q ∧ (i2' + len + j2'.toUInt64).toNat = b' * 2 ^ (q + 2) + j2' + 2 ^ (q + 1) ∧ (i2' + len + j2'.toUInt64 + s).toNat = b' * 2 ^ (q + 2) + j2' + 2 ^ (q + 1) + 2 ^ q ∧ (i2' + j2'.toUInt64).toNat < m ∧ (i2' + j2'.toUInt64 + s).toNat < m ∧ (i2' + len + j2'.toUInt64).toNat < m ∧ (i2' + len + j2'.toUInt64 + s).toNat < m
Position arithmetic for block b', with ℕ size bound m.
theorem
radix4Middle_comp
{N : ℕ}
(inverse : Bool)
(roots : Vector UInt32 N)
(s len k1 k2 b_start : ℕ)
(a : Vector UInt32 N)
:
radix4Middle inverse roots s len (k1 + k2) b_start a = radix4Middle inverse roots s len k2 (b_start + k1) (radix4Middle inverse roots s len k1 b_start a)
theorem
radix4_block_ne_pos
(q b j2 j2nat : ℕ)
(hj2 : j2 < 2 ^ q)
(hj2_lt : j2nat < 2 ^ q)
(hj2_ne : j2 ≠ j2nat)
(posval : ℕ)
(hpos :
posval = b * 2 ^ (q + 2) + j2nat ∨ posval = b * 2 ^ (q + 2) + j2nat + 2 ^ q ∨ posval = b * 2 ^ (q + 2) + j2nat + 2 ^ (q + 1) ∨ posval = b * 2 ^ (q + 2) + j2nat + 2 ^ (q + 1) + 2 ^ q)
:
Butterfly at j2 ≠ j2nat does not touch any of the four block positions for j2nat. Isolated so omega runs in minimal context (avoids slow hypothesis scanning).