improve NTT and add tests

This commit is contained in:
Balazs Komuves 2026-06-13 13:03:06 +02:00
parent 7dc831f0d7
commit 1116e2791a
No known key found for this signature in database
GPG Key ID: F63B7AEF18435562
4 changed files with 109 additions and 9 deletions

View File

@ -126,6 +126,10 @@ func squareFp* (y: Fp): Fp = ( var z : Fp = y ; square(z) ; return z )
func squareFp2*(y: Fp2): Fp2 = ( var z : Fp2 = y ; square(z) ; return z )
func squareFr* (y: Fr): Fr = ( var z : Fr = y ; square(z) ; return z )
func `/`*(x, y: Fp ): Fp = ( var z : Fp = x ; z *= invFr(y) ; return z )
func `/`*(x, y: Fp2): Fp2 = ( var z : Fp2 = x ; z *= invFr(y) ; return z )
func `/`*(x, y: Fr ): Fr = ( var z : Fr = x ; z *= invFr(y) ; return z )
# template/generic instantiation of `pow_vartime` from here
# /Users/bkomuves/.nimble/pkgs/constantine-0.0.1/constantine/math/arithmetic/finite_fields.nim(389, 7) template/generic instantiation of `fieldMod` from here
# /Users/bkomuves/.nimble/pkgs/constantine-0.0.1/constantine/math/config/curves_prop_field_derived.nim(67, 5) Error: undeclared identifier: 'getCurveOrder'

View File

@ -18,7 +18,7 @@ import groth16/math/domain
func forwardNTT_worker( m: int
, srcStride: int
, gpows: seq[Fr[BN254_Snarks]]
, src: seq[Fr[BN254_Snarks]] , srcOfs: int
, src: seq[Fr[BN254_Snarks]] , srcOfs: int
, buf: var seq[Fr[BN254_Snarks]] , bufOfs: int
, tgt: var seq[Fr[BN254_Snarks]] , tgtOfs: int ) =
case m
@ -93,12 +93,11 @@ func extendAndForwardNTT*(src: seq[Fr[BN254_Snarks]], D: Domain): seq[Fr[BN254_S
#-------------------------------------------------------------------------------
const oneHalfFr* = fromHex(Fr[BN254_Snarks], "0x183227397098d014dc2822db40c0ac2e9419f4243cdcb848a1f0fac9f8000001")
# unscaled!
func inverseNTT_worker( m: int
, tgtStride: int
, gpows: seq[Fr[BN254_Snarks]]
, src: seq[Fr[BN254_Snarks]] , srcOfs: int
, src: seq[Fr[BN254_Snarks]] , srcOfs: int
, buf: var seq[Fr[BN254_Snarks]] , bufOfs: int
, tgt: var seq[Fr[BN254_Snarks]] , tgtOfs: int ) =
case m
@ -109,8 +108,6 @@ func inverseNTT_worker( m: int
of 1:
tgt[tgtOfs ] = ( src[srcOfs] + src[srcOfs+1] )
tgt[tgtOfs+tgtStride] = ( src[srcOfs] - src[srcOfs+1] )
div2( tgt[tgtOfs ] )
div2( tgt[tgtOfs+tgtStride] )
else:
let N : int = 1 shl m
@ -119,7 +116,6 @@ func inverseNTT_worker( m: int
for j in 0..<halfN:
buf[bufOfs+j ] = ( src[srcOfs+j] + src[srcOfs+j+halfN] )
buf[bufOfs+j+halfN] = ( src[srcOfs+j] - src[srcOfs+j+halfN] ) * gpows[ j*tgtStride ]
div2( buf[bufOfs+j ] )
inverseNTT_worker( m-1
, tgtStride shl 1
@ -137,7 +133,7 @@ func inverseNTT_worker( m: int
#---------------------------------------
# inverse number-theoretical transform (corresponds to polynomial interpolation)
func inverseNTT*(src: seq[Fr[BN254_Snarks]], D: Domain): seq[Fr[BN254_Snarks]] =
func toggableInverseNTT(src: seq[Fr[BN254_Snarks]], D: Domain, do_rescale: bool): seq[Fr[BN254_Snarks]] =
assert( D.domainSize == (1 shl D.logDomainSize) , "domain must have a power-of-two size" )
assert( D.domainSize == src.len , "input must have the same size as the domain" )
var buf = newSeq[Fr[BN254_Snarks]]( 2 * D.domainSize )
@ -147,7 +143,7 @@ func inverseNTT*(src: seq[Fr[BN254_Snarks]], D: Domain): seq[Fr[BN254_Snarks]] =
let N = D.domainSize
let halFN = N div 2
var gpows = newSeq[Fr[BN254_Snarks]]( halFN )
var x = oneHalfFr
var x = oneFr
let ginv = invFr( D.domainGen )
for i in 0..<halfN:
gpows[i] = x
@ -159,6 +155,24 @@ func inverseNTT*(src: seq[Fr[BN254_Snarks]], D: Domain): seq[Fr[BN254_Snarks]] =
, src , 0
, buf , 0
, tgt , 0 )
if do_rescale:
var invN : Fr[BN254_Snarks]
invN.fromInt(N)
invN.inv()
for i in 0..<N:
tgt[i] *= invN
return tgt
#-------------------------------------------------------------------------------
# inverse number-theoretical transform (corresponds to polynomial interpolation)
func inverseNTT*(src: seq[Fr[BN254_Snarks]], D: Domain): seq[Fr[BN254_Snarks]] =
toggableInverseNTT(src, D, true)
# inverse number-theoretical transform, without the 1/N rescaling
func unscaledInverseNTT*(src: seq[Fr[BN254_Snarks]], D: Domain): seq[Fr[BN254_Snarks]] =
toggableInverseNTT(src, D, false)
#-------------------------------------------------------------------------------

81
tests/groth16/testFFT.nim Normal file
View File

@ -0,0 +1,81 @@
{.used.}
import std/unittest
import constantine/math/arithmetic
import constantine/named/properties_fields
import groth16/bn128/fields
import groth16/bn128/curves
import groth16/bn128/rnd
import groth16/bn128/debug
import groth16/math/domain
import groth16/math/ntt
#-------------------------------------------------------------------------------
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 === ys)
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
#-------------------------------------------------------------------------------
suite "FFT checks":
suite "field NTT":
let D = createDomain(128)
let N = D.domainSize
let xs : seq[Fr[BN254_Snarks]] = randFrSeq(N)
test "INTT(NTT(xs) == xs":
let ys = forwardNTT(xs, D)
let zs = inverseNTT(ys, D)
check isEqualFrSeq( xs , zs )
test "NTT(INTT(xs) == xs":
let ys = inverseNTT(xs, D)
let zs = forwardNTT(ys, D)
check isEqualFrSeq( xs , zs )
test "unscaledINTT(NTT(xs) == N * xs":
let ys = forwardNTT(xs, D)
let zs = unscaledInverseNTT(ys, D)
var ws = xs
scaleFrSeqInPlace(intToFr(N) , ws)
check isEqualFrSeq( ws , zs )
#-------------------------------------------------------------------------------

View File

@ -1,4 +1,5 @@
import ./groth16/testFFT
import ./groth16/testPtCompression
import ./groth16/testCurve
import ./groth16/testProver