Browse Source

fft benchmarking

x 3 năm trước cách đây
mục cha
commit
9376389beb

+ 44 - 6
script/research/zk/fft/fft5.sage → script/research/zk/fft/benchmark-fft.sage

@@ -1,4 +1,5 @@
-import itertools
+import itertools, time
+from tabulate import tabulate
 
 
 def find_ext_order(p, n):
 def find_ext_order(p, n):
     N = 1
     N = 1
@@ -93,7 +94,7 @@ def test1():
 def random_test():
 def random_test():
     p = random_prime(1000)
     p = random_prime(1000)
     #n = 16
     #n = 16
-    n = int(2^ZZ.random_element(2, 10))
+    n = 2^8
     assert p.is_prime()
     assert p.is_prime()
     N = find_ext_order(p, n)
     N = find_ext_order(p, n)
     print(f"p = {p}")
     print(f"p = {p}")
@@ -117,12 +118,49 @@ def random_test():
 
 
     ω_powers = vector(ω^i for i in range(n/2))
     ω_powers = vector(ω^i for i in range(n/2))
     fT = vectorify(X, f, n)
     fT = vectorify(X, f, n)
+
+    start = time.time()
     dft = calc_dft(n, ω_powers, fT)
     dft = calc_dft(n, ω_powers, fT)
-    print()
-    print(f"DFT(f) = {dft}")
+    dft_duration = time.time() - start
+
+    print(f"DFT time: {dft_duration}")
+
+    start = time.time()
     f_evals = [f(X=ω^i) for i in range(n)]
     f_evals = [f(X=ω^i) for i in range(n)]
-    print(f"f(ω^i) = {f_evals}")
+    eval_duration = time.time() - start
+
+    print(f"Eval time: {eval_duration}")
+
+    #print()
+    #print(f"DFT(f) = {dft}")
+    #print()
+    #print(f"f(ω^i) = {f_evals}")
+    assert dft == f_evals
+
+    return dft_duration, eval_duration
+
+def timing_info():
+    table = []
+    total_dft, total_eval = 0, 0
+    success = 0
+    for i in range(20):
+        print(f"Trial: {i}")
+        try:
+            dft, eval = random_test()
+        except AssertionError:
+            table.append((i, "Error", "Error"))
+            continue
+        table.append((i, dft, eval))
+        total_dft += dft
+        total_eval += eval
+        success += 1
+    avg_dft = total_dft / success
+    avg_eval = total_eval / success
+    table.append(("", "", ""))
+    table.append(("Average:", avg_dft, avg_eval))
+    print(tabulate(table, headers=["#", "DFT", "Naive"]))
 
 
 #test1()
 #test1()
-random_test()
+#timing_info()
+#random_test()