mirror of
https://github.com/logos-storage/nim-groth16.git
synced 2026-07-21 07:59:32 +00:00
improve NTT and add tests
This commit is contained in:
parent
7dc831f0d7
commit
1116e2791a
@ -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'
|
||||
|
||||
@ -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
81
tests/groth16/testFFT.nim
Normal 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 )
|
||||
|
||||
#-------------------------------------------------------------------------------
|
||||
|
||||
|
||||
@ -1,4 +1,5 @@
|
||||
|
||||
import ./groth16/testFFT
|
||||
import ./groth16/testPtCompression
|
||||
import ./groth16/testCurve
|
||||
import ./groth16/testProver
|
||||
|
||||
Loading…
x
Reference in New Issue
Block a user