add some quick&dirty FFT benchmarks

This commit is contained in:
Balazs Komuves 2026-06-13 14:20:17 +02:00
parent 2231d800c7
commit f17c0389b9
No known key found for this signature in database
GPG Key ID: F63B7AEF18435562
8 changed files with 213 additions and 47 deletions

2
bench/.gitignore vendored Normal file
View File

@ -0,0 +1,2 @@
.DS_Store
bench_fft

104
bench/bench_fft.nim Normal file
View File

@ -0,0 +1,104 @@
#
# nimble build -d:release
#
import strformat
import taskpools
import constantine/math/arithmetic
import constantine/named/properties_fields
import groth16/bn128
import groth16/bn128/msm
import groth16/bn128/arrays
import groth16/math/domain
import groth16/math/ntt
import groth16/dynamic/group_fft
import shared
#-------------------------------------------------------------------------------
when isMainModule:
echo "quick & dirty FFT benchmarks"
let nthreads: int = 8
let N: int = 8192
let D = createDomain(N)
let xs : seq[Fr[BN254_Snarks]] = randFrSeq(N)
let gs : seq[AffG1] = randG1Seq(N)
var ys, zs: seq[Fr[BN254_Snarks]]
var hs, rs: seq[AffG1]
var pool = Taskpool.new(nthreads)
#-----------------------------------------------------------------------------
echo "\n----------------------------------------"
echo "*** scalar multiplications\n"
hs = newSeq[AffG1](N)
withMeasureTime(true , fmt"{N} individual scalar multiplications "):
for i in 0..<N:
hs[i] = xs[i] ** gs[i]
var sum_naive: AffG1
withMeasureTime(true , fmt"naive simulated MSM of size {N} "):
for i in 0..<N:
hs[i] = xs[i] ** gs[i]
if i==0:
sum_naive = hs[i]
else:
sum_naive += hs[i]
var sum_msm: AffG1
withMeasureTime(true , fmt"proper MSM of size {N} "):
sum_msm = msmConstantineG1( xs, gs )
var sum_multi: AffG1
withMeasureTime(true , fmt"multithreaded MSM of size {N} ({nthreads} threads) "):
let sum_multi = msmMultiThreadedG1( xs , gs , pool )
echo "naive == msm : " & $(sum_naive === sum_msm)
#-----------------------------------------------------------------------------
echo "\n----------------------------------------"
echo "*** field FFTs\n"
withMeasureTime(true , fmt"field NTT of size {N} "):
ys = forwardNTT(xs, D)
withMeasureTime(true , fmt"unscaled field INTT of size {N} "):
zs = unscaledInverseNTT(ys, D)
withMeasureTime(true , fmt"field INTT of size {N} "):
zs = inverseNTT(ys, D)
echo "xs == zs : " & $isEqualFrSeq(xs , zs)
#-----------------------------------------------------------------------------
echo "\n----------------------------------------"
echo "*** group FFTs\n"
withMeasureTime(true , fmt"group FFT of size {N} "):
hs = forwardGroupFFT(gs, D)
withMeasureTime(true , fmt"unscaled group IFFT of size {N} "):
rs = unscaledInverseGroupFFT(hs, D)
withMeasureTime(true , fmt"group IFFT of size {N} "):
rs = inverseGroupFFT(hs, D)
echo "gs == rs : " & $isEqualG1Seq(gs , rs)
echo "\n"
#-------------------------------------------------------------------------------

13
bench/bench_fft.nimble Normal file
View File

@ -0,0 +1,13 @@
version = "0.0.1"
author = "Balazs Komuves"
description = "FFT benchmarks"
license = "MIT OR Apache-2.0"
binDir = "build"
bin = @["bench_fft"]
requires "nim >= 2.2.0"
requires "https://github.com/status-im/nim-taskpools >= 0.0.5"
requires "https://github.com/mratsim/constantine"
requires "groth16 >= 0.1.2"

1
bench/nim.cfg Normal file
View File

@ -0,0 +1 @@
--path:".."

23
bench/shared.nim Normal file
View File

@ -0,0 +1,23 @@
import strformat
import times, strutils
#-------------------------------------------------------------------------------
func seconds*(x: float): string = fmt"{x:.4f} seconds"
func quoted*(s: string): string = fmt"`{s:s}`"
template withMeasureTime*(doPrint: bool, text: string, code: untyped) =
block:
if doPrint:
let t0 = epochTime()
code
let elapsed = epochTime() - t0
let elapsedStr = elapsed.formatFloat(format = ffDecimal, precision = 4)
echo ( text & " took " & elapsedStr & " seconds" )
else:
code
#-------------------------------------------------------------------------------

View File

@ -4,7 +4,7 @@ author = "Balazs Komuves"
description = "Groth16 proof system"
license = "MIT OR Apache-2.0"
skipDirs = @["groth16/example"]
skipDirs = @["groth16/example","bench"]
binDir = "build"
namedBin = {"cli/cli_main": "nim-groth16"}.toTable()
installExt = @["nim"]

61
groth16/bn128/arrays.nim Normal file
View File

@ -0,0 +1,61 @@
import constantine/math/arithmetic
import constantine/named/properties_fields
import groth16/bn128/fields
import groth16/bn128/curves
import groth16/bn128/rnd
#-------------------------------------------------------------------------------
# Fr arrays
proc randFrSeq*(N: int) : seq[Fr[BN254_Snarks]] =
var xs : seq[Fr[BN254_Snarks]] = newSeq[Fr[BN254_Snarks]]( N )
for i in 0..<N:
xs[i] = randFr()
return xs
func isEqualFrSeq*(xs: seq[Fr[BN254_Snarks]], ys: seq[Fr[BN254_Snarks]]): bool =
let N = xs.len
let M = ys.len
if N != M:
return false
else:
var ok: bool = true
for i in 0..<N:
ok = ok and (xs[i] === ys[i])
return ok
proc scaleFrSeqInPlace*(s: Fr[BN254_Snarks], arr: var seq[Fr[BN254_Snarks]] ) =
let N = arr.len
for i in 0..<N:
arr[i] *= s
#-------------------------------------------------------------------------------
# G1 arrays
proc randG1Seq*(N: int) : seq[G1] =
var xs : seq[G1] = newSeq[G1]( N )
for i in 0..<N:
xs[i] = randG1()
return xs
func isEqualG1Seq*(xs: seq[G1], ys: seq[G1] ): bool =
let N = xs.len
let M = ys.len
if N != M:
return false
else:
var ok: bool = true
for i in 0..<N:
ok = ok and (xs[i] === ys[i])
return ok
proc scaleG1SeqInPlace*(s: Fr[BN254_Snarks], arr: var seq[G1] ) =
let N = arr.len
for i in 0..<N:
arr[i] = s ** arr[i]
#-------------------------------------------------------------------------------

View File

@ -11,6 +11,7 @@ import groth16/bn128/fields
import groth16/bn128/curves
import groth16/bn128/rnd
import groth16/bn128/debug
import groth16/bn128/arrays
import groth16/math/domain
import groth16/math/ntt
@ -19,51 +20,6 @@ import groth16/dynamic/group_fft
#-------------------------------------------------------------------------------
proc randFrSeq(N: int) : seq[Fr[BN254_Snarks]] =
var xs : seq[Fr[BN254_Snarks]] = newSeq[Fr[BN254_Snarks]]( N )
for i in 0..<N:
xs[i] = randFr()
return xs
proc scaleFrSeqInPlace(s: Fr[BN254_Snarks], arr: var seq[Fr[BN254_Snarks]] ) =
let N = arr.len
for i in 0..<N:
arr[i] *= s
func isEqualFrSeq(xs: seq[Fr[BN254_Snarks]], ys: seq[Fr[BN254_Snarks]]): bool =
let N = xs.len
let M = ys.len
if N != M:
return false
else:
var ok: bool = true
for i in 0..<N:
ok = ok and (xs[i] === ys[i])
return ok
#---------------------------------------
proc randG1Seq(N: int) : seq[G1] =
var xs : seq[G1] = newSeq[G1]( N )
for i in 0..<N:
xs[i] = randG1()
return xs
func isEqualG1Seq(xs: seq[G1], ys: seq[G1] ): bool =
let N = xs.len
let M = ys.len
if N != M:
return false
else:
var ok: bool = true
for i in 0..<N:
ok = ok and (xs[i] === ys[i])
return ok
#-------------------------------------------------------------------------------
suite "field NTT checks":
let D = createDomain(128)
@ -100,11 +56,17 @@ suite "group FFT checks (for the group G1)":
let rs = inverseGroupFFT(hs, D)
check isEqualG1Seq( gs , rs )
test "NTT(INTT(xs) == xs":
test "FFT(IFFT(xs) == xs":
let hs = inverseGroupFFT(gs, D)
let rs = forwardGroupFFT(hs, D)
check isEqualG1Seq( gs , rs )
test "FFT(unscaledIFFT(gs) == N * gs":
let hs = unscaledInverseGroupFFT(gs, D)
let rs = forwardGroupFFT(hs, D)
var qs = gs
scaleG1SeqInPlace(intToFr(N) , qs)
check isEqualG1Seq( qs , rs )
#-------------------------------------------------------------------------------