benchmark-fft.sage 6.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240
  1. import itertools, time
  2. from tabulate import tabulate
  3. def find_ext_order(p, n):
  4. N = 1
  5. while True:
  6. pNx_order = p^N - 1
  7. # Does n divide the group order 𝔽_{p^N}^×?
  8. if pNx_order % n == 0:
  9. return N
  10. N += 1
  11. def find_nth_root_unity(K, p, N, n):
  12. # It cannot be a quadratic residue if n is odd
  13. #assert n % 2 == 1
  14. # So there is an nth root of unity in p^N. Now we have to find it.
  15. pNx_order = p^N - 1
  16. assert n > 1
  17. assert int(pNx_order) % n == 0
  18. ω = K.multiplicative_generator()
  19. assert ω^pNx_order == 1
  20. ω = ω^(pNx_order/n)
  21. assert ω^n == 1
  22. assert ω^(n - 1) != 1
  23. return ω
  24. def vectorify(X, f, n):
  25. assert f.degree() < n
  26. fT = vector([f[i] for i in range(f.degree() + 1)] +
  27. # Zero padding
  28. [0 for _ in range(n - f.degree() - 1)])
  29. assert len(fT) == n
  30. # Just check decomposed polynomial is in the correct order
  31. assert sum([fT[i]*X^i for i in range(n)]) == f
  32. return fT
  33. def dot(a, b):
  34. assert len(a) == len(b)
  35. return [a_i*b_i for a_i, b_i in zip(a, b)]
  36. # ABC, DEF -> ADBECF
  37. def alternate(list1, list2):
  38. return itertools.chain(*zip(list1, list2))
  39. def calc_dft(n, ω_powers, f):
  40. m = len(f)
  41. if m == 1:
  42. return f
  43. g, h = vector(f[:m/2]), vector(f[m/2:])
  44. r = g + h
  45. s = dot(g - h, ω_powers)
  46. ω_powers = vector(ω_i for ω_i in ω_powers[::2])
  47. rT = calc_dft(n, ω_powers, r)
  48. sT = calc_dft(n, ω_powers, s)
  49. result = list(alternate(rT, sT))
  50. return result
  51. def test1():
  52. p = 199
  53. #n = 16
  54. n = 8
  55. assert p.is_prime()
  56. N = find_ext_order(p, n)
  57. print(f"p = {p}")
  58. print(f"n = {n}")
  59. print(f"N = {N}")
  60. print(f"p^N = {p^N}")
  61. K.<a> = GF(p^N, repr="int")
  62. ω = find_nth_root_unity(K, p, N, n)
  63. print()
  64. L.<X> = K[]
  65. #f = 9*X^7 + 45*X^6 + 33*X^5 + 7*X^3 + X^2 + 110*X + 4
  66. f = 7*X^3 + X^2 + 110*X + 4
  67. assert f.degree() < n/2
  68. print(f"f = {f}")
  69. print()
  70. ω_powers = vector(ω^i for i in range(n/2))
  71. fT = vectorify(X, f, n)
  72. dft = calc_dft(n, ω_powers, fT)
  73. print()
  74. print(f"DFT(f) = {dft}")
  75. f_evals = [f(X=ω^i) for i in range(n)]
  76. print(f"f(ω^i) = {f_evals}")
  77. def random_test():
  78. p = random_prime(1000)
  79. #d = int(ZZ.random_element(6, 8))
  80. #n = 2^d
  81. n = 256
  82. assert p.is_prime()
  83. N = find_ext_order(p, n)
  84. print(f"p = {p}")
  85. print(f"n = {n}")
  86. print(f"N = {N}")
  87. print(f"p^N = {p^N}")
  88. K.<a> = GF(p^N, repr="int")
  89. ω = find_nth_root_unity(K, p, N, n)
  90. print(f"ω = {ω}")
  91. L.<X> = K[]
  92. #f = 9*X^7 + 45*X^6 + 33*X^5 + 7*X^3 + X^2 + 110*X + 4
  93. f = 0
  94. for i in range(n/2):
  95. f += ZZ.random_element(0, 200) * X^i
  96. assert f.degree() < n/2
  97. #print(f"f = {f}")
  98. ω_powers = vector(ω^i for i in range(n/2))
  99. fT = vectorify(X, f, n)
  100. start = time.time()
  101. dft = calc_dft(n, ω_powers, fT)
  102. dft_duration = time.time() - start
  103. print(f"DFT time: {dft_duration}")
  104. start = time.time()
  105. f_evals = [f(X=ω^i) for i in range(n)]
  106. eval_duration = time.time() - start
  107. print(f"Eval time: {eval_duration}")
  108. print()
  109. #print()
  110. #print(f"DFT(f) = {dft}")
  111. #print()
  112. #print(f"f(ω^i) = {f_evals}")
  113. assert dft == f_evals
  114. return dft_duration, eval_duration, n, log(p).n()
  115. def timing_info():
  116. table = []
  117. total_dft, total_eval = 0, 0
  118. success = 0
  119. for i in range(20):
  120. print(f"Trial: {i}")
  121. try:
  122. dft, eval, n, log_p = random_test()
  123. except AssertionError:
  124. table.append((i, "Error", "", "", ""))
  125. continue
  126. table.append((i, f"{dft:.5f}", f"{eval:.5f}", n, log_p))
  127. total_dft += dft
  128. total_eval += eval
  129. success += 1
  130. avg_dft = total_dft / success
  131. avg_eval = total_eval / success
  132. table.append(("", "", ""))
  133. table.append(("Average:", f"{avg_dft:.5f}", f"{avg_eval:.5f}"))
  134. print(tabulate(table, headers=["#", "DFT", "Naive", "n", "log_p"]))
  135. def test_root_of_unity():
  136. p = random_prime(1000)
  137. d = int(ZZ.random_element(2, 8))
  138. n = 2^d
  139. assert p.is_prime()
  140. N = find_ext_order(p, n)
  141. print(f"p = {p}")
  142. print(f"n = {n}")
  143. print(f"N = {N}")
  144. print(f"p^N = {p^N}")
  145. K.<a> = GF(p^N, repr="int")
  146. ω = find_nth_root_unity(K, p, N, n)
  147. print(f"ω = {ω}")
  148. print()
  149. def pallas_base_test():
  150. print("Pallas base Fp test")
  151. d = 11
  152. print(f"n = 2^{d}")
  153. p = 28948022309329048855892746252171976963363056481941560715954676764349967630337
  154. n = 2^d
  155. K = GF(p, repr="int")
  156. ω = K(5)^((p - 1)/n)
  157. assert ω^n == 1
  158. assert ω^(n - 1) != 1
  159. ω_powers = vector(ω^i for i in range(n/2))
  160. table = []
  161. number_trials = 10
  162. total_dft_duration, total_naive_duration = 0, 0
  163. for trial_i in range(number_trials):
  164. print(f"Trial: {trial_i}")
  165. print("Generating random polynomial...")
  166. fT = ([K.random_element() for _ in range(n/2)]
  167. + [K(0) for _ in range(n/2)])
  168. assert len(fT) == n
  169. start = time.time()
  170. dft = calc_dft(n, ω_powers, fT)
  171. dft_duration = time.time() - start
  172. print(f"DFT time: {dft_duration}")
  173. def eval_poly(fT, x):
  174. accum = 0
  175. current_x = 1
  176. for a in fT:
  177. accum += a*current_x
  178. current_x *= x
  179. return accum
  180. start = time.time()
  181. evals = [eval_poly(fT, ω^i) for i in range(n)]
  182. eval_duration = time.time() - start
  183. print(f"Eval time: {eval_duration}")
  184. print()
  185. table.append((trial_i, f"{dft_duration:.5f}", f"{eval_duration:.5f}"))
  186. total_dft_duration += dft_duration
  187. total_naive_duration += eval_duration
  188. avg_dft = total_dft_duration / number_trials
  189. avg_naive = total_naive_duration / number_trials
  190. table.append(("", "", ""))
  191. table.append(("Average:", f"{avg_dft:.5f}", f"{avg_naive:.5f}"))
  192. print(tabulate(table, headers=["#", "DFT", "Naive"]))
  193. pallas_base_test()
  194. #test1()
  195. #timing_info()
  196. #random_test()
  197. #for i in range(50):
  198. # test_root_of_unity()