fft.sage 1.5 KB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061626364
  1. # See also https://cp-algorithms.com/algebra/fft.html
  2. q = 0x40000000000000000000000000000000224698fc0994a8dd8c46eb2100000001
  3. K = GF(q)
  4. P.<X> = K[]
  5. def get_omega():
  6. generator = K(5)
  7. assert (q - 1) % 2^32 == 0
  8. # Root of unity
  9. t = (q - 1) / 2^32
  10. omega = generator**t
  11. assert omega != 1
  12. assert omega^(2^16) != 1
  13. assert omega^(2^31) != 1
  14. assert omega^(2^32) == 1
  15. return omega
  16. # Order of this element is 2^32
  17. omega = get_omega()
  18. k = 3
  19. n = 2^k
  20. omega = omega^(2^32 / n)
  21. assert omega^n == 1
  22. f = 6*X^7 + 7*X^5 + 3*X^2 + X
  23. def fft(F):
  24. print(f"fft({F})")
  25. # On the first invocation:
  26. #assert len(F) == n
  27. N = len(F)
  28. if N == 1:
  29. print(" returning 1")
  30. return F
  31. omega_prime = omega^(n/N)
  32. assert omega_prime^(n - 1) != 1
  33. assert omega_prime^N == 1
  34. # Split into even and odd powers of X
  35. F_e = [a for a in F[::2]]
  36. print(" Evens:", F_e)
  37. F_o = [a for a in F[1::2]]
  38. print(" Odds:", F_o)
  39. y_e, y_o = fft(F_e), fft(F_o)
  40. print(f"y_e = {y_e}, y_o = {y_o}")
  41. y = [0] * N
  42. for j in range(N / 2):
  43. y[j] = y_e[j] + omega_prime^j * y_o[j]
  44. y[j + N / 2] = y_e[j] - omega_prime^j * y_o[j]
  45. print(f" returning y = {y}")
  46. return y
  47. print("f =", f)
  48. evals = fft(list(f))
  49. print("evals =", evals)
  50. print("{omega^i : i in {0, 1, ..., n - 1}} =", [omega^i for i in range(n)])
  51. evals2 = [f(omega^i) for i in range(n)]
  52. print("{f(omega^i) for all omega^i} =", evals2)
  53. assert evals == evals2