fft4.sage 2.4 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103
  1. import itertools
  2. def find_ext_order(p, n):
  3. N = 1
  4. while True:
  5. pNx_order = p^N - 1
  6. # Does n divide the group order 𝔽_{p^N}^×?
  7. if pNx_order % n == 0:
  8. return N
  9. N += 1
  10. def find_nth_root_unity(K, p, N, n):
  11. # It cannot be a quadratic residue if n is odd
  12. #assert n % 2 == 1
  13. # So there is an nth root of unity in p^N. Now we have to find it.
  14. pNx_order = p^N - 1
  15. ω = K.multiplicative_generator()
  16. ω = ω^(pNx_order/n)
  17. assert ω^n == 1
  18. assert ω^(n - 1) != 1
  19. return ω
  20. def vectorify(X, f, n):
  21. assert f.degree() < n
  22. fT = vector([f[i] for i in range(f.degree() + 1)] +
  23. # Zero padding
  24. [0 for _ in range(n - f.degree() - 1)])
  25. assert len(fT) == n
  26. # Just check decomposed polynomial is in the correct order
  27. assert sum([fT[i]*X^i for i in range(n)]) == f
  28. return fT
  29. def dot(a, b):
  30. assert len(a) == len(b)
  31. return [a_i*b_i for a_i, b_i in zip(a, b)]
  32. # ABC, DEF -> ADBECF
  33. def alternate(list1, list2):
  34. return itertools.chain(*zip(list1, list2))
  35. def calc_dft(n, ω_powers, f):
  36. m = len(f)
  37. indent = " " * (n - m)
  38. print(f"{indent}calc_dft({ω_powers}, {f})")
  39. print(f"{indent} m = {m}")
  40. if m == 1:
  41. print(f"{indent} m = 1 so return f")
  42. return f
  43. g, h = vector(f[:m/2]), vector(f[m/2:])
  44. print(f"{indent} g = {g}")
  45. print(f"{indent} h = {h}")
  46. r = g + h
  47. s = dot(g - h, ω_powers)
  48. print(f"{indent} r = {r}")
  49. print(f"{indent} s = {s}")
  50. print()
  51. ω_powers = vector(ω_i for ω_i in ω_powers[::2])
  52. rT = calc_dft(n, ω_powers, r)
  53. sT = calc_dft(n, ω_powers, s)
  54. result = list(alternate(rT, sT))
  55. print(f"{indent}return {result}")
  56. return result
  57. def test():
  58. p = 199
  59. #n = 16
  60. n = 8
  61. assert p.is_prime()
  62. N = find_ext_order(p, n)
  63. print(f"p = {p}")
  64. print(f"n = {n}")
  65. print(f"N = {N}")
  66. print(f"p^N = {p^N}")
  67. K.<a> = GF(p^N, repr="int")
  68. ω = find_nth_root_unity(K, p, N, n)
  69. print(f"ω = {ω}")
  70. print()
  71. L.<X> = K[]
  72. #f = 9*X^7 + 45*X^6 + 33*X^5 + 7*X^3 + X^2 + 110*X + 4
  73. f = 7*X^3 + X^2 + 110*X + 4
  74. assert f.degree() < n/2
  75. print(f"f = {f}")
  76. print()
  77. ω_powers = vector(ω^i for i in range(n/2))
  78. fT = vectorify(X, f, n)
  79. dft = calc_dft(n, ω_powers, fT)
  80. print()
  81. print(f"DFT(f) = {dft}")
  82. f_evals = [f(X=ω^i) for i in range(n)]
  83. print(f"f(ω^i) = {f_evals}")
  84. test()