fft2.sage 3.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131
  1. # sage: w.multiplicative_order()
  2. # 1330
  3. # sage: 11^3 - 1
  4. # 1330
  5. # sage: K.<w> = GF(11^3, repr="int")
  6. # sage: w
  7. # 11
  8. p = 199
  9. # sage: factor(199^3 - 1)
  10. # 2 * 3^3 * 11 * 13267
  11. n = 3^3
  12. n = 10
  13. n = 5
  14. assert p.is_prime()
  15. def find_ext_order(p, n):
  16. N = 1
  17. while True:
  18. pNx_order = p^N - 1
  19. # Does n divide the group order 𝔽_{p^N}^×?
  20. if pNx_order % n == 0:
  21. return N
  22. N += 1
  23. def find_nth_root_unity(K, n):
  24. # It cannot be a quadratic residue if n is odd
  25. #assert n % 2 == 1
  26. # So there is an nth root of unity in p^N. Now we have to find it.
  27. pNx_order = p^N - 1
  28. ω = K.multiplicative_generator()
  29. ω = ω^(pNx_order/n)
  30. assert ω^n == 1
  31. assert ω^(n - 1) != 1
  32. return ω
  33. N = find_ext_order(p, n)
  34. print(f"N = {N}")
  35. print()
  36. K.<a> = GF(p^N, repr="int")
  37. ω = find_nth_root_unity(K, n)
  38. L.<X> = K[]
  39. f = 3*X^4 + 7*X^3 + X^2 + 4
  40. g = 2*X^4 + 2*X^2 + 110
  41. f = X^2 + 2*X + 4
  42. g = 2*X^2 + 110
  43. assert f.degree() < n/2
  44. assert g.degree() < n/2
  45. assert f.degree() + g.degree() < n
  46. print(f"f = {f}")
  47. print(f"g = {g}")
  48. print(f"fg = {f*g}")
  49. print()
  50. def vectorify(f):
  51. assert f.degree() < n
  52. fT = vector([f[i] for i in range(f.degree() + 1)] +
  53. # Zero padding
  54. [0 for _ in range(n - f.degree() - 1)])
  55. assert len(fT) == n
  56. # Just check decomposed polynomial is in the correct order
  57. assert sum([fT[i]*X^i for i in range(n)]) == f
  58. return fT
  59. fT = vectorify(f)
  60. gT = vectorify(g)
  61. def nXn_vandermonde(n, ω):
  62. # We hardcode this one so you know what is looks like
  63. if n == 5:
  64. Vω = matrix([
  65. [1, 1, 1, 1, 1],
  66. [1, ω^1, ω^2, ω^3, ω^4],
  67. [1, ω^2, ω^4, ω^1, ω^3],
  68. [1, ω^3, ω^1, ω^4, ω^2],
  69. [1, ω^4, ω^3, ω^2, ω^1],
  70. ])
  71. return Vω
  72. # This is the code to generate it
  73. Vω = matrix([[ω^(i * j) for j in range(n)] for i in range(n)])
  74. return Vω
  75. Vω = nXn_vandermonde(n, ω)
  76. Vω_inv = nXn_vandermonde(n, ω^-1)/n
  77. # Lemma: V_ω^{-1} = 1/n V_{ω^-1}
  78. assert Vω^-1 == Vω_inv
  79. DFT_ω_f = Vω * fT
  80. f_evals = [f(X=ω^i) for i in range(n)]
  81. print(f"DFT_ω(f) = {DFT_ω_f}")
  82. print(f"f(ω^i) = {f_evals}")
  83. print()
  84. DFT_ω_g = Vω * gT
  85. g_evals = [g(X=ω^i) for i in range(n)]
  86. print(f"DFT_ω(g) = {DFT_ω_g}")
  87. print(f"g(ω^i) = {g_evals}")
  88. print()
  89. def convolution(f, g):
  90. return f*g % (X^n - 1)
  91. def pointwise_prod(fT, gT):
  92. return [a_i*b_i for a_i, b_i in zip(fT, gT)]
  93. print(f"deg(f) + deg(g) = {f.degree() + g.degree()}")
  94. fжg = convolution(f, g)
  95. print(f"f☼g = {fжg}")
  96. assert fжg == f*g
  97. fжgT = vectorify(fжg)
  98. DFT_ω_fжg = Vω * fжgT
  99. for i in range(n):
  100. assert fжg(X=ω^i) == f(ω^i)*g(ω^i)
  101. print(f"DFT_ω(f☼g) = {DFT_ω_fжg}")
  102. DFT_fg_prod = vector(pointwise_prod(DFT_ω_f, DFT_ω_g))
  103. print(f"DFT_ω(f)·DFT_ω(g) = {DFT_fg_prod}")
  104. assert DFT_ω_fжg == DFT_fg_prod
  105. inv_DFT_fg = Vω_inv * DFT_fg_prod
  106. fgT = vectorify(f*g)
  107. assert inv_DFT_fg == fgT
  108. print(f"DFT^-1(DFT_ω(f)·DFT_ω(g)) = {inv_DFT_fg}")