From 1116e2791a314755fdbbbfb8a66bb265c16f793b Mon Sep 17 00:00:00 2001 From: Balazs Komuves Date: Sat, 13 Jun 2026 13:03:06 +0200 Subject: [PATCH] improve NTT and add tests --- groth16/bn128/fields.nim | 4 ++ groth16/math/ntt.nim | 32 +++++++++++----- tests/groth16/testFFT.nim | 81 +++++++++++++++++++++++++++++++++++++++ tests/test.nim | 1 + 4 files changed, 109 insertions(+), 9 deletions(-) create mode 100644 tests/groth16/testFFT.nim diff --git a/groth16/bn128/fields.nim b/groth16/bn128/fields.nim index 1a560c5..9fe6a8f 100644 --- a/groth16/bn128/fields.nim +++ b/groth16/bn128/fields.nim @@ -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' diff --git a/groth16/math/ntt.nim b/groth16/math/ntt.nim index b69a8ab..7018f02 100644 --- a/groth16/math/ntt.nim +++ b/groth16/math/ntt.nim @@ -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..