From a12d64f85cc38469f49869ca0eb292292ef1e7fe Mon Sep 17 00:00:00 2001 From: Arya Tabaie Date: Tue, 22 Sep 2026 13:46:35 -0500 Subject: [PATCH 01/18] feat: multilin funcs for koalabear.E6 Signed-off-by: Arya Tabaie --- internal/generator/main.go | 10 ++++ internal/generator/polynomial/generate.go | 22 +++++++ .../polynomial/template/doc.multilin.go.tmpl | 2 + .../polynomial/template/multilin.test.go.tmpl | 58 +++++++++++++++++-- 4 files changed, 86 insertions(+), 6 deletions(-) create mode 100644 internal/generator/polynomial/template/doc.multilin.go.tmpl diff --git a/internal/generator/main.go b/internal/generator/main.go index 8846a9e080..de8d0ca99e 100644 --- a/internal/generator/main.go +++ b/internal/generator/main.go @@ -83,6 +83,16 @@ func main() { field.WithExtensions(), field.WithIOP(), )) + + // multilinear polynomials over the degree-6 extension E6 + if f.Name == "koalabear" { + extInfo := fieldConfig.FieldDependency{ + FieldPackagePath: "github.com/consensys/gnark-crypto/field/" + f.Name + "/extensions", + FieldPackageName: "extensions", + ElementType: "extensions.E6", + } + assertNoError(polynomial.GenerateMultilin(extInfo, filepath.Join(outputDir, "extensions", "polynomial"), true, gen)) + } }(conf) } diff --git a/internal/generator/polynomial/generate.go b/internal/generator/polynomial/generate.go index 7739d6b4ea..cfb0bdec70 100644 --- a/internal/generator/polynomial/generate.go +++ b/internal/generator/polynomial/generate.go @@ -28,3 +28,25 @@ func Generate(conf config.FieldDependency, baseDir string, generateTests bool, g polyGen := common.NewDefaultGenerator(template.FS) return polyGen.Generate(conf, "polynomial", "", "", entries...) } + +// GenerateMultilin generates only the multilinear polynomial API (MultiLin and +// its memory Pool) over conf.ElementType. Unlike Generate, it does not require +// the element type to have the full base-field API (e.g. SetInt64, Vector +// helpers, big.Int conversions), so it can be used for field extensions such +// as koalabear's E6. +func GenerateMultilin(conf config.FieldDependency, baseDir string, generateTests bool, gen *common.Generator) error { + entries := []bavard.Entry{ + {File: filepath.Join(baseDir, "doc.go"), Templates: []string{"doc.multilin.go.tmpl"}}, + {File: filepath.Join(baseDir, "multilin.go"), Templates: []string{"multilin.go.tmpl"}}, + {File: filepath.Join(baseDir, "pool.go"), Templates: []string{"pool.go.tmpl"}}, + } + + if generateTests { + entries = append(entries, + bavard.Entry{File: filepath.Join(baseDir, "multilin_test.go"), Templates: []string{"multilin.test.go.tmpl"}}, + ) + } + + polyGen := common.NewDefaultGenerator(template.FS) + return polyGen.Generate(conf, "polynomial", "", "", entries...) +} diff --git a/internal/generator/polynomial/template/doc.multilin.go.tmpl b/internal/generator/polynomial/template/doc.multilin.go.tmpl new file mode 100644 index 0000000000..cf777371e9 --- /dev/null +++ b/internal/generator/polynomial/template/doc.multilin.go.tmpl @@ -0,0 +1,2 @@ +// Package polynomial provides multilinear polynomial methods over {{.ElementType}}. +package polynomial diff --git a/internal/generator/polynomial/template/multilin.test.go.tmpl b/internal/generator/polynomial/template/multilin.test.go.tmpl index 2bef046b30..c9b0aa68b7 100644 --- a/internal/generator/polynomial/template/multilin.test.go.tmpl +++ b/internal/generator/polynomial/template/multilin.test.go.tmpl @@ -11,7 +11,9 @@ func TestFoldBilinear(t *testing.T) { // f = c₀ + c₁ X₁ + c₂ X₂ + c₃ X₁ X₂ var coefficients [4]{{.ElementType}} - fr.Vector(coefficients[:]).MustSetRandom() + for i := range coefficients { + coefficients[i].MustSetRandom() + } var r {{.ElementType}} r.MustSetRandom() @@ -47,9 +49,12 @@ func TestFoldBilinear(t *testing.T) { // TODO: Benchmark folding? Algorithms is pretty straightforward; unless we want to measure how well memory management is working func TestFoldedEqTable(t *testing.T) { + var one {{.ElementType}} + one.SetOne() + q := make([]{{.ElementType}}, 2) - q[0].SetInt64(2) - q[1].SetInt64(3) + q[0].Double(&one) // q₀ = 2 + q[1].Add(&q[0], &one) // q₁ = 3 m := make(MultiLin, 4) m[0].SetOne() @@ -58,9 +63,6 @@ func TestFoldedEqTable(t *testing.T) { eq := make([]{{.ElementType}}, 4) p := make([]{{.ElementType}}, 2) - var one {{.ElementType}} - one.SetOne() - for p0 := range 2 { p[1].SetZero() for p1 := range 2 { @@ -75,3 +77,47 @@ func TestFoldedEqTable(t *testing.T) { } } + +func TestEvaluateMatchesEqInnerProduct(t *testing.T) { + const n = 3 + + var one {{.ElementType}} + one.SetOne() + + m := make(MultiLin, 1<>(n-1-j))&1 == 1 { + b[j] = one + } else { + b[j].SetZero() + } + } + term = EvalEq(r, b) + term.Mul(&term, &m[i]) + expected.Add(&expected, &term) + } + + pool := NewPool(1 << n) + for _, p := range []*Pool{nil, &pool} { + clone := m.Clone() + res := m.Evaluate(r, p) + assert.True(t, res.Equal(&expected), "evaluation disagrees with Eq inner product") + for i := range m { + assert.True(t, m[i].Equal(&clone[i]), "Evaluate must not modify its receiver") + } + } + assert.Equal(t, n, m.NumVars()) +} From 7e86a97411ab814fa659007a9dc6e449adf43ed6 Mon Sep 17 00:00:00 2001 From: Arya Tabaie Date: Thu, 24 Sep 2026 08:56:52 -0500 Subject: [PATCH 02/18] feat: more multilin functions for koala E6 Signed-off-by: Arya Tabaie --- field/babybear/extensions/e6.go | 33 +++ field/koalabear/extensions/e6.go | 33 +++ field/koalabear/extensions/e6_vector_test.go | 13 + field/koalabear/extensions/polynomial/doc.go | 7 + .../extensions/polynomial/multilin.go | 179 +++++++++++++ .../extensions/polynomial/multilin_test.go | 85 ++++++ .../extensions/polynomial/polynomial.go | 253 ++++++++++++++++++ .../extensions/polynomial/polynomial_test.go | 245 +++++++++++++++++ field/koalabear/extensions/polynomial/pool.go | 191 +++++++++++++ .../generator/field/config/field_config.go | 1 + .../field/template/extensions/e6.go.tmpl | 33 +++ internal/generator/main.go | 6 +- internal/generator/polynomial/generate.go | 22 -- .../polynomial/template/doc.multilin.go.tmpl | 2 - .../polynomial/template/multilin.test.go.tmpl | 62 +---- .../polynomial/template/polynomial.go.tmpl | 6 + .../template/polynomial.test.go.tmpl | 6 + 17 files changed, 1099 insertions(+), 78 deletions(-) create mode 100644 field/koalabear/extensions/polynomial/doc.go create mode 100644 field/koalabear/extensions/polynomial/multilin.go create mode 100644 field/koalabear/extensions/polynomial/multilin_test.go create mode 100644 field/koalabear/extensions/polynomial/polynomial.go create mode 100644 field/koalabear/extensions/polynomial/polynomial_test.go create mode 100644 field/koalabear/extensions/polynomial/pool.go delete mode 100644 internal/generator/polynomial/template/doc.multilin.go.tmpl diff --git a/field/babybear/extensions/e6.go b/field/babybear/extensions/e6.go index f831a4d207..afbfa69237 100644 --- a/field/babybear/extensions/e6.go +++ b/field/babybear/extensions/e6.go @@ -72,6 +72,20 @@ func (z *E6) SetOne() *E6 { return z } +// SetInt64 sets z to v (embedded via the unique ring homomorphism ℤ→𝔽₆) and returns z +func (z *E6) SetInt64(v int64) *E6 { + *z = E6{} + z.B0.A0.SetInt64(v) + return z +} + +// SetUint64 sets z to v (embedded via the unique ring homomorphism ℤ→𝔽₆) and returns z +func (z *E6) SetUint64(v uint64) *E6 { + *z = E6{} + z.B0.A0.SetUint64(v) + return z +} + // MulByElement multiplies an element in E6 by an element in fr. // y may alias a coordinate of x, so we copy it first. func (z *E6) MulByElement(x *E6, y *fr.Element) *E6 { @@ -470,6 +484,25 @@ func ButterflyE6(a, b *E6) { // VectorE6 represents a vector of E6 elements type VectorE6 []E6 +// SetRandom sets all elements of vector to random values, returning the first error encountered, if any. +func (vector VectorE6) SetRandom() error { + for i := range vector { + if _, err := vector[i].SetRandom(); err != nil { + return err + } + } + return nil +} + +// MustSetRandom sets all elements of vector to random values, panicking if an error is encountered. +func (vector VectorE6) MustSetRandom() { + for i := range vector { + if _, err := vector[i].SetRandom(); err != nil { + panic(err) + } + } +} + // Butterfly computes the in-place butterfly operation on two vectors of E6 elements. // If other overlaps with vector, the result is undefined; the caller should use a temp vector. func (vector VectorE6) Butterfly(other VectorE6) { diff --git a/field/koalabear/extensions/e6.go b/field/koalabear/extensions/e6.go index 5ad81b19fe..964d524620 100644 --- a/field/koalabear/extensions/e6.go +++ b/field/koalabear/extensions/e6.go @@ -73,6 +73,20 @@ func (z *E6) SetOne() *E6 { return z } +// SetInt64 sets z to v (embedded via the unique ring homomorphism ℤ→𝔽₆) and returns z +func (z *E6) SetInt64(v int64) *E6 { + *z = E6{} + z.B0.A0.SetInt64(v) + return z +} + +// SetUint64 sets z to v (embedded via the unique ring homomorphism ℤ→𝔽₆) and returns z +func (z *E6) SetUint64(v uint64) *E6 { + *z = E6{} + z.B0.A0.SetUint64(v) + return z +} + // MulByElement multiplies an element in E6 by an element in fr. // y may alias a coordinate of x, so we copy it first. func (z *E6) MulByElement(x *E6, y *fr.Element) *E6 { @@ -471,6 +485,25 @@ func ButterflyE6(a, b *E6) { // VectorE6 represents a vector of E6 elements type VectorE6 []E6 +// SetRandom sets all elements of vector to random values, returning the first error encountered, if any. +func (vector VectorE6) SetRandom() error { + for i := range vector { + if _, err := vector[i].SetRandom(); err != nil { + return err + } + } + return nil +} + +// MustSetRandom sets all elements of vector to random values, panicking if an error is encountered. +func (vector VectorE6) MustSetRandom() { + for i := range vector { + if _, err := vector[i].SetRandom(); err != nil { + panic(err) + } + } +} + // Butterfly computes the in-place butterfly operation on two vectors of E6 elements. // If other overlaps with vector, the result is undefined; the caller should use a temp vector. func (vector VectorE6) Butterfly(other VectorE6) { diff --git a/field/koalabear/extensions/e6_vector_test.go b/field/koalabear/extensions/e6_vector_test.go index c71fa8cd74..e2b08e8e2e 100644 --- a/field/koalabear/extensions/e6_vector_test.go +++ b/field/koalabear/extensions/e6_vector_test.go @@ -11,6 +11,19 @@ import ( "github.com/stretchr/testify/require" ) +func TestVectorE6MustSetRandom(t *testing.T) { + vector := make(VectorE6, 17) + vector.MustSetRandom() + + var zero E6 + for i := range vector { + require.False(t, vector[i].Equal(&zero), "element %d was left zero", i) + for j := range vector[:i] { + require.False(t, vector[i].Equal(&vector[j]), "elements %d and %d collided", j, i) + } + } +} + func TestVectorE6ButterflyPair(t *testing.T) { for _, size := range []int{2, 4, 8, 16, 18} { t.Run(strconv.Itoa(size), func(t *testing.T) { diff --git a/field/koalabear/extensions/polynomial/doc.go b/field/koalabear/extensions/polynomial/doc.go new file mode 100644 index 0000000000..aa346f3ea3 --- /dev/null +++ b/field/koalabear/extensions/polynomial/doc.go @@ -0,0 +1,7 @@ +// Copyright 2020-2026 Consensys Software Inc. +// Licensed under the Apache License, Version 2.0. See the LICENSE file for details. + +// Code generated by consensys/gnark-crypto DO NOT EDIT + +// Package polynomial provides polynomial methods and commitment schemes. +package polynomial diff --git a/field/koalabear/extensions/polynomial/multilin.go b/field/koalabear/extensions/polynomial/multilin.go new file mode 100644 index 0000000000..d0c45379fa --- /dev/null +++ b/field/koalabear/extensions/polynomial/multilin.go @@ -0,0 +1,179 @@ +// Copyright 2020-2026 Consensys Software Inc. +// Licensed under the Apache License, Version 2.0. See the LICENSE file for details. + +// Code generated by consensys/gnark-crypto DO NOT EDIT + +package polynomial + +import ( + "math/bits" + + "github.com/consensys/gnark-crypto/field/koalabear/extensions" + "github.com/consensys/gnark-crypto/utils" +) + +// MultiLin tracks the values of a (dense i.e. not sparse) multilinear polynomial +// The variables are X₁ through Xₙ where n = log(len(.)) +// .[∑ᵢ 2ⁱ⁻¹ bₙ₋ᵢ] = the polynomial evaluated at (b₁, b₂, ..., bₙ) +// It is understood that any hypercube evaluation can be extrapolated to a multilinear polynomial +type MultiLin []extensions.E6 + +// Fold is partial evaluation function k[X₁, X₂, ..., Xₙ] → k[X₂, ..., Xₙ] by setting X₁=r +func (m *MultiLin) Fold(r extensions.E6) { + mid := len(*m) / 2 + + bottom, top := (*m)[:mid], (*m)[mid:] + + var t extensions.E6 // no need to update the top part + + // updating bookkeeping table + // knowing that the polynomial f ∈ (k[X₂, ..., Xₙ])[X₁] is linear, we would get f(r) = f(0) + r(f(1) - f(0)) + // the following loop computes the evaluations of f(r) accordingly: + // f(r, b₂, ..., bₙ) = f(0, b₂, ..., bₙ) + r(f(1, b₂, ..., bₙ) - f(0, b₂, ..., bₙ)) + for i := range mid { + // table[i] ← table[i] + r (table[i + mid] - table[i]) + t.Sub(&top[i], &bottom[i]) + t.Mul(&t, &r) + bottom[i].Add(&bottom[i], &t) + } + + *m = (*m)[:mid] +} + +func (m *MultiLin) FoldParallel(r extensions.E6) utils.Task { + mid := len(*m) / 2 + bottom, top := (*m)[:mid], (*m)[mid:] + + *m = bottom + + return func(start, end int) { + var t extensions.E6 // no need to update the top part + for i := start; i < end; i++ { + // table[i] ← table[i] + r (table[i + mid] - table[i]) + t.Sub(&top[i], &bottom[i]) + t.Mul(&t, &r) + bottom[i].Add(&bottom[i], &t) + } + } +} + +func (m MultiLin) Sum() extensions.E6 { + s := m[0] + for i := 1; i < len(m); i++ { + s.Add(&s, &m[i]) + } + return s +} + +func _clone(m MultiLin, p *Pool) MultiLin { + if p == nil { + return m.Clone() + } else { + return p.Clone(m) + } +} + +func _dump(m MultiLin, p *Pool) { + if p != nil { + p.Dump(m) + } +} + +// Evaluate extrapolate the value of the multilinear polynomial corresponding to m +// on the given coordinates +func (m MultiLin) Evaluate(coordinates []extensions.E6, p *Pool) extensions.E6 { + // Folding is a mutating operation + bkCopy := _clone(m, p) + + // Evaluate step by step through repeated folding (i.e. evaluation at the first remaining variable) + for _, r := range coordinates { + bkCopy.Fold(r) + } + + result := bkCopy[0] + + _dump(bkCopy, p) + return result +} + +// Clone creates a deep copy of a bookkeeping table. +// Both multilinear interpolation and sumcheck require folding an underlying +// array, but folding changes the array. To do both one requires a deep copy +// of the bookkeeping table. +func (m MultiLin) Clone() MultiLin { + res := make(MultiLin, len(m)) + copy(res, m) + return res +} + +// Add two bookKeepingTables +func (m *MultiLin) Add(left, right MultiLin) { + size := len(left) + // Check that left and right have the same size + if len(right) != size || len(*m) != size { + panic("left, right and destination must have the right size") + } + + // Add elementwise + for i := range size { + (*m)[i].Add(&left[i], &right[i]) + } +} + +// EvalEq computes Eq(q₁, ... , qₙ, h₁, ... , hₙ) = Π₁ⁿ Eq(qᵢ, hᵢ) +// where Eq(x,y) = xy + (1-x)(1-y) = 1 - x - y + xy + xy interpolates +// +// _________________ +// | | | +// | 0 | 1 | +// |_______|_______| +// y | | | +// | 1 | 0 | +// |_______|_______| +// +// x +// +// In other words the polynomial evaluated here is the multilinear extrapolation of +// one that evaluates to q' == h' for vectors q', h' of binary values +func EvalEq(q, h []extensions.E6) extensions.E6 { + var res, nxt, one, sum extensions.E6 + one.SetOne() + for i := range len(q) { + nxt.Mul(&q[i], &h[i]) // nxt <- qᵢ * hᵢ + nxt.Double(&nxt) // nxt <- 2 * qᵢ * hᵢ + nxt.Add(&nxt, &one) // nxt <- 1 + 2 * qᵢ * hᵢ + sum.Add(&q[i], &h[i]) // sum <- qᵢ + hᵢ TODO: Why not subtract one by one from nxt? More parallel? + + if i == 0 { + res.Sub(&nxt, &sum) // nxt <- 1 + 2 * qᵢ * hᵢ - qᵢ - hᵢ + } else { + nxt.Sub(&nxt, &sum) // nxt <- 1 + 2 * qᵢ * hᵢ - qᵢ - hᵢ + res.Mul(&res, &nxt) // res <- res * nxt + } + } + return res +} + +// Eq sets m to the representation of the polynomial Eq(q₁, ..., qₙ, *, ..., *) × m[0] +func (m *MultiLin) Eq(q []extensions.E6) { + n := len(q) + + if len(*m) != 1<= 0; i-- { + res.Mul(&res, v) + res.Add(&res, &(*p)[i]) + } + + return res +} + +// Clone returns a copy of the polynomial +func (p *Polynomial) Clone() Polynomial { + _p := make(Polynomial, len(*p)) + copy(_p, *p) + return _p +} + +// Set to another polynomial +func (p *Polynomial) Set(p1 Polynomial) { + if len(*p) != len(p1) { + *p = p1.Clone() + return + } + + for i := range len(p1) { + (*p)[i].Set(&p1[i]) + } +} + +// AddConstantInPlace adds a constant to the polynomial, modifying p +func (p *Polynomial) AddConstantInPlace(c *extensions.E6) { + for i := range len(*p) { + (*p)[i].Add(&(*p)[i], c) + } +} + +// SubConstantInPlace subs a constant to the polynomial, modifying p +func (p *Polynomial) SubConstantInPlace(c *extensions.E6) { + for i := range len(*p) { + (*p)[i].Sub(&(*p)[i], c) + } +} + +// ScaleInPlace multiplies p by v, modifying p +func (p *Polynomial) ScaleInPlace(c *extensions.E6) { + for i := range len(*p) { + (*p)[i].Mul(&(*p)[i], c) + } +} + +// Scale multiplies p0 by v, storing the result in p +func (p *Polynomial) Scale(c *extensions.E6, p0 Polynomial) { + if len(*p) != len(p0) { + *p = make(Polynomial, len(p0)) + } + for i := range len(p0) { + (*p)[i].Mul(c, &p0[i]) + } +} + +// Add adds p1 to p2 +// This function allocates a new slice unless p == p1 or p == p2 +func (p *Polynomial) Add(p1, p2 Polynomial) *Polynomial { + + bigger := p1 + smaller := p2 + if len(bigger) < len(smaller) { + bigger, smaller = smaller, bigger + } + + if len(*p) == len(bigger) && (&(*p)[0] == &bigger[0]) { + for i := range len(smaller) { + (*p)[i].Add(&(*p)[i], &smaller[i]) + } + return p + } + + if len(*p) == len(smaller) && (&(*p)[0] == &smaller[0]) { + for i := range len(smaller) { + (*p)[i].Add(&(*p)[i], &bigger[i]) + } + *p = append(*p, bigger[len(smaller):]...) + return p + } + + res := make(Polynomial, len(bigger)) + copy(res, bigger) + for i := range len(smaller) { + res[i].Add(&res[i], &smaller[i]) + } + *p = res + return p +} + +// Sub subtracts p2 from p1 +// TODO make interface more consistent with Add +func (p *Polynomial) Sub(p1, p2 Polynomial) *Polynomial { + if len(p1) != len(p2) || len(p2) != len(*p) { + return nil + } + for i := range len(*p) { + (*p)[i].Sub(&p1[i], &p2[i]) + } + return p +} + +// Equal checks equality between two polynomials +func (p *Polynomial) Equal(p1 Polynomial) bool { + if (*p == nil) != (p1 == nil) { + return false + } + + if len(*p) != len(p1) { + return false + } + + for i := range p1 { + if !(*p)[i].Equal(&p1[i]) { + return false + } + } + + return true +} + +func (p Polynomial) SetZero() { + for i := range len(p) { + p[i].SetZero() + } +} + +// InterpolateOnRange maps vector v to polynomial f +// such that f(i) = v[i] for 0 ≤ i < len(v). +// len(f) = len(v) and deg(f) ≤ len(v) - 1 +func InterpolateOnRange(v []extensions.E6) Polynomial { + nEvals := uint8(len(v)) + if int(nEvals) != len(v) { + panic("interpolation method too inefficient for nEvals > 255") + } + lagrange := getLagrangeBasis(nEvals) + + var res Polynomial + res.Scale(&v[0], lagrange[0]) + + temp := make(Polynomial, nEvals) + + for i := uint8(1); i < nEvals; i++ { + temp.Scale(&v[i], lagrange[i]) + res.Add(res, temp) + } + + return res +} + +// lagrange bases used by InterpolateOnRange +var lagrangeBasis sync.Map + +func getLagrangeBasis(domainSize uint8) []Polynomial { + if res, ok := lagrangeBasis.Load(domainSize); ok { + return res.([]Polynomial) + } + + // not found. compute + var res []Polynomial + if domainSize >= 2 { + res = computeLagrangeBasis(domainSize) + } else if domainSize == 1 { + res = []Polynomial{make(Polynomial, 1)} + res[0][0].SetOne() + } + lagrangeBasis.Store(domainSize, res) + + return res +} + +// computeLagrangeBasis precomputes in explicit coefficient form for each 0 ≤ l < domainSize the polynomial +// pₗ := X (X-1) ... (X-l-1) (X-l+1) ... (X - domainSize + 1) / ( l (l-1) ... 2 (-1) ... (l - domainSize +1) ) +// Note that pₗ(l) = 1 and pₗ(n) = 0 if 0 ≤ l < domainSize, n ≠ l +func computeLagrangeBasis(domainSize uint8) []Polynomial { + + constTerms := make([]extensions.E6, domainSize) + for i := range domainSize { + constTerms[i].SetInt64(-int64(i)) + } + + res := make([]Polynomial, domainSize) + multScratch := make(Polynomial, domainSize-1) + + // compute pₗ + for l := range domainSize { + + // TODO @Tabaie Optimize this with some trees? O(log(domainSize)) polynomial mults instead of O(domainSize)? Then again it would be fewer big poly mults vs many small poly mults + d := uint8(0) //d is the current degree of res + for i := range domainSize { + if i == l { + continue + } + if d == 0 { + res[l] = make(Polynomial, domainSize) + res[l][domainSize-2] = constTerms[i] + res[l][domainSize-1].SetOne() + } else { + current := res[l][domainSize-d-2:] + timesConst := multScratch[domainSize-d-2:] + + timesConst.Scale(&constTerms[i], current[1:]) //TODO: Directly double and add since constTerms are tiny? (even less than 4 bits) + nonLeading := current[0 : d+1] + + nonLeading.Add(nonLeading, timesConst) + + } + d++ + } + + } + + // We have pₗ(i≠l)=0. Now scale so that pₗ(l)=1 + // Replace the constTerms with norms + for l := range domainSize { + constTerms[l].Neg(&constTerms[l]) + constTerms[l] = res[l].Eval(&constTerms[l]) + } + constTerms = extensions.BatchInvertE6(constTerms) + for l := range domainSize { + res[l].ScaleInPlace(&constTerms[l]) + } + + return res +} diff --git a/field/koalabear/extensions/polynomial/polynomial_test.go b/field/koalabear/extensions/polynomial/polynomial_test.go new file mode 100644 index 0000000000..711574f4ec --- /dev/null +++ b/field/koalabear/extensions/polynomial/polynomial_test.go @@ -0,0 +1,245 @@ +// Copyright 2020-2026 Consensys Software Inc. +// Licensed under the Apache License, Version 2.0. See the LICENSE file for details. + +// Code generated by consensys/gnark-crypto DO NOT EDIT + +package polynomial + +import ( + "math/big" + "testing" + + "github.com/consensys/gnark-crypto/field/koalabear/extensions" + "github.com/leanovate/gopter" + "github.com/leanovate/gopter/gen" + "github.com/leanovate/gopter/prop" + "github.com/stretchr/testify/assert" +) + +func TestPolynomialEval(t *testing.T) { + + // build polynomial + f := make(Polynomial, 20) + for i := range 20 { + f[i].SetOne() + } + + // random value + var point extensions.E6 + point.MustSetRandom() + + // compute manually f(val) + var expectedEval, one, den extensions.E6 + var expo big.Int + one.SetOne() + expo.SetUint64(20) + expectedEval.Exp(point, &expo). + Sub(&expectedEval, &one) + den.Sub(&point, &one) + expectedEval.Div(&expectedEval, &den) + + // compute purported evaluation + purportedEval := f.Eval(&point) + + // check + if !purportedEval.Equal(&expectedEval) { + t.Fatal("polynomial evaluation failed") + } +} + +func TestPolynomialAddConstantInPlace(t *testing.T) { + + // build polynomial + f := make(Polynomial, 20) + for i := range 20 { + f[i].SetOne() + } + + // constant to add + var c extensions.E6 + c.MustSetRandom() + + // add constant + f.AddConstantInPlace(&c) + + // check + var expectedCoeffs, one extensions.E6 + one.SetOne() + expectedCoeffs.Add(&one, &c) + for i := range 20 { + if !f[i].Equal(&expectedCoeffs) { + t.Fatal("AddConstantInPlace failed") + } + } +} + +func TestPolynomialSubConstantInPlace(t *testing.T) { + + // build polynomial + f := make(Polynomial, 20) + for i := range 20 { + f[i].SetOne() + } + + // constant to sub + var c extensions.E6 + c.MustSetRandom() + + // sub constant + f.SubConstantInPlace(&c) + + // check + var expectedCoeffs, one extensions.E6 + one.SetOne() + expectedCoeffs.Sub(&one, &c) + for i := range 20 { + if !f[i].Equal(&expectedCoeffs) { + t.Fatal("SubConstantInPlace failed") + } + } +} + +func TestPolynomialScaleInPlace(t *testing.T) { + + // build polynomial + f := make(Polynomial, 20) + for i := range 20 { + f[i].SetOne() + } + + // constant to scale by + var c extensions.E6 + c.MustSetRandom() + + // scale by constant + f.ScaleInPlace(&c) + + // check + for i := range 20 { + if !f[i].Equal(&c) { + t.Fatal("ScaleInPlace failed") + } + } + +} + +func TestPolynomialAdd(t *testing.T) { + + // build unbalanced polynomials + f1 := make(Polynomial, 20) + f1Backup := make(Polynomial, 20) + for i := range 20 { + f1[i].SetOne() + f1Backup[i].SetOne() + } + f2 := make(Polynomial, 10) + f2Backup := make(Polynomial, 10) + for i := range 10 { + f2[i].SetOne() + f2Backup[i].SetOne() + } + + // expected result + var one, two extensions.E6 + one.SetOne() + two.Double(&one) + expectedSum := make(Polynomial, 20) + for i := range 10 { + expectedSum[i].Set(&two) + } + for i := 10; i < 20; i++ { + expectedSum[i].Set(&one) + } + + // caller is empty + var g Polynomial + g.Add(f1, f2) + if !g.Equal(expectedSum) { + t.Fatal("add polynomials fails") + } + if !f1.Equal(f1Backup) { + t.Fatal("side effect, f1 should not have been modified") + } + if !f2.Equal(f2Backup) { + t.Fatal("side effect, f2 should not have been modified") + } + + // all operands are distinct + _f1 := f1.Clone() + _f1.Add(f1, f2) + if !_f1.Equal(expectedSum) { + t.Fatal("add polynomials fails") + } + if !f1.Equal(f1Backup) { + t.Fatal("side effect, f1 should not have been modified") + } + if !f2.Equal(f2Backup) { + t.Fatal("side effect, f2 should not have been modified") + } + + // first operand = caller + _f1 = f1.Clone() + _f2 := f2.Clone() + _f1.Add(_f1, _f2) + if !_f1.Equal(expectedSum) { + t.Fatal("add polynomials fails") + } + if !_f2.Equal(f2Backup) { + t.Fatal("side effect, _f2 should not have been modified") + } + + // second operand = caller + _f1 = f1.Clone() + _f2 = f2.Clone() + _f1.Add(_f2, _f1) + if !_f1.Equal(expectedSum) { + t.Fatal("add polynomials fails") + } + if !_f2.Equal(f2Backup) { + t.Fatal("side effect, _f2 should not have been modified") + } +} + +func TestPrecomputeLagrange(t *testing.T) { + + testForDomainSize := func(domainSize uint8) bool { + polys := computeLagrangeBasis(domainSize) + + for l := range domainSize { + for i := range domainSize { + var I extensions.E6 + I.SetUint64(uint64(i)) + y := polys[l].Eval(&I) + + if i == l && !y.IsOne() || i != l && !y.IsZero() { + t.Errorf("domainSize = %d: p_%d(%d) = %s", domainSize, l, i, y.String()) + return false + } + } + } + return true + } + + t.Parallel() + parameters := gopter.DefaultTestParameters() + + const maxLagrangeDomainSize = 12 + + parameters.MinSuccessfulTests = maxLagrangeDomainSize + + properties := gopter.NewProperties(parameters) + + properties.Property("l'th lagrange polynomials must evaluate to 1 on l and 0 on other values in the domain", prop.ForAll( + testForDomainSize, + gen.UInt8Range(2, maxLagrangeDomainSize), + )) + + properties.TestingRun(t, gopter.ConsoleReporter(false)) +} + +func TestLagrangeCache(t *testing.T) { + for _, i := range []int{5, 2, 8, 4, 6, 3, 0} { + b := getLagrangeBasis(uint8(i)) + assert.Equal(t, b, getLagrangeBasis(uint8(i))) // second call must yield the same result + } +} diff --git a/field/koalabear/extensions/polynomial/pool.go b/field/koalabear/extensions/polynomial/pool.go new file mode 100644 index 0000000000..3f246202f7 --- /dev/null +++ b/field/koalabear/extensions/polynomial/pool.go @@ -0,0 +1,191 @@ +// Copyright 2020-2026 Consensys Software Inc. +// Licensed under the Apache License, Version 2.0. See the LICENSE file for details. + +// Code generated by consensys/gnark-crypto DO NOT EDIT + +package polynomial + +import ( + "encoding/json" + "fmt" + "runtime" + "sort" + "sync" + "unsafe" + + "github.com/consensys/gnark-crypto/field/koalabear/extensions" +) + +// Memory management for polynomials +// WARNING: This is not thread safe TODO: Make sure that is not a problem +// TODO: There is a lot of "unsafe" memory management here and needs to be vetted thoroughly + +type sizedPool struct { + maxN int + pool sync.Pool + stats poolStats +} + +type inUseData struct { + allocatedFor []uintptr + pool *sizedPool +} + +type Pool struct { + //lock sync.Mutex + inUse sync.Map + subPools []sizedPool +} + +func (p *sizedPool) get(n int) *extensions.E6 { + p.stats.make(n) + return p.pool.Get().(*extensions.E6) +} + +func (p *sizedPool) put(ptr *extensions.E6) { + p.stats.dump() + p.pool.Put(ptr) +} + +func NewPool(maxN ...int) (pool Pool) { + + sort.Ints(maxN) + pool = Pool{ + subPools: make([]sizedPool, len(maxN)), + } + + for i := range pool.subPools { + subPool := &pool.subPools[i] + subPool.maxN = maxN[i] + subPool.pool = sync.Pool{ + New: func() any { + subPool.stats.Allocated++ + return getDataPointer(make([]extensions.E6, 0, subPool.maxN)) + }, + } + } + return +} + +func (p *Pool) findCorrespondingPool(n int) *sizedPool { + poolI := 0 + for poolI < len(p.subPools) && n > p.subPools[poolI].maxN { + poolI++ + } + return &p.subPools[poolI] // out of bounds error here would mean that n is too large +} + +func (p *Pool) Make(n int) []extensions.E6 { + pool := p.findCorrespondingPool(n) + ptr := pool.get(n) + p.addInUse(ptr, pool) + return unsafe.Slice(ptr, n) +} + +// Dump dumps a set of polynomials into the pool +func (p *Pool) Dump(slices ...[]extensions.E6) { + for _, slice := range slices { + ptr := getDataPointer(slice) + if metadata, ok := p.inUse.Load(ptr); ok { + p.inUse.Delete(ptr) + metadata.(inUseData).pool.put(ptr) + } else { + panic("attempting to dump a slice not created by the pool") + } + } +} + +func (p *Pool) addInUse(ptr *extensions.E6, pool *sizedPool) { + pcs := make([]uintptr, 2) + n := runtime.Callers(3, pcs) + + if prevPcs, ok := p.inUse.Load(ptr); ok { // TODO: remove if unnecessary for security + panic(fmt.Errorf("re-allocated non-dumped slice, previously allocated at %v", runtime.CallersFrames(prevPcs.(inUseData).allocatedFor))) + } + p.inUse.Store(ptr, inUseData{ + allocatedFor: pcs[:n], + pool: pool, + }) +} + +func printFrame(frame runtime.Frame) { + fmt.Printf("\t%s line %d, function %s\n", frame.File, frame.Line, frame.Function) +} + +func (p *Pool) printInUse() { + fmt.Println("slices never dumped allocated at:") + p.inUse.Range(func(_, pcs any) bool { + fmt.Println("-------------------------") + + var frame runtime.Frame + frames := runtime.CallersFrames(pcs.(inUseData).allocatedFor) + more := true + for more { + frame, more = frames.Next() + printFrame(frame) + } + return true + }) +} + +type poolStats struct { + Used int + Allocated int + ReuseRate float64 + InUse int + GreatestNUsed int + SmallestNUsed int +} + +type poolsStats struct { + SubPools []poolStats + InUse int +} + +func (s *poolStats) make(n int) { + s.Used++ + s.InUse++ + if n > s.GreatestNUsed { + s.GreatestNUsed = n + } + if s.SmallestNUsed == 0 || s.SmallestNUsed > n { + s.SmallestNUsed = n + } +} + +func (s *poolStats) dump() { + s.InUse-- +} + +func (s *poolStats) finalize() { + s.ReuseRate = float64(s.Used) / float64(s.Allocated) +} + +func getDataPointer(slice []extensions.E6) *extensions.E6 { + return (*extensions.E6)(unsafe.SliceData(slice)) +} + +func (p *Pool) PrintPoolStats() { + InUse := 0 + subStats := make([]poolStats, len(p.subPools)) + for i := range p.subPools { + subPool := &p.subPools[i] + subPool.stats.finalize() + subStats[i] = subPool.stats + InUse += subPool.stats.InUse + } + + stats := poolsStats{ + SubPools: subStats, + InUse: InUse, + } + serialized, _ := json.MarshalIndent(stats, "", " ") + fmt.Println(string(serialized)) + p.printInUse() +} + +func (p *Pool) Clone(slice []extensions.E6) []extensions.E6 { + res := p.Make(len(slice)) + copy(res, slice) + return res +} diff --git a/internal/generator/field/config/field_config.go b/internal/generator/field/config/field_config.go index 4870453cb6..3b175b99e2 100644 --- a/internal/generator/field/config/field_config.go +++ b/internal/generator/field/config/field_config.go @@ -831,4 +831,5 @@ type FieldDependency struct { ElementType string FieldPackagePath string FieldPackageName string + ExtensionDegree int } diff --git a/internal/generator/field/template/extensions/e6.go.tmpl b/internal/generator/field/template/extensions/e6.go.tmpl index 6962306d11..0101c58129 100644 --- a/internal/generator/field/template/extensions/e6.go.tmpl +++ b/internal/generator/field/template/extensions/e6.go.tmpl @@ -68,6 +68,20 @@ func (z *E6) SetOne() *E6 { return z } +// SetInt64 sets z to v (embedded via the unique ring homomorphism ℤ→𝔽₆) and returns z +func (z *E6) SetInt64(v int64) *E6 { + *z = E6{} + z.B0.A0.SetInt64(v) + return z +} + +// SetUint64 sets z to v (embedded via the unique ring homomorphism ℤ→𝔽₆) and returns z +func (z *E6) SetUint64(v uint64) *E6 { + *z = E6{} + z.B0.A0.SetUint64(v) + return z +} + // MulByElement multiplies an element in E6 by an element in fr. // y may alias a coordinate of x, so we copy it first. func (z *E6) MulByElement(x *E6, y *fr.Element) *E6 { @@ -466,6 +480,25 @@ func ButterflyE6(a, b *E6) { // VectorE6 represents a vector of E6 elements type VectorE6 []E6 +// SetRandom sets all elements of vector to random values, returning the first error encountered, if any. +func (vector VectorE6) SetRandom() error { + for i := range vector { + if _, err := vector[i].SetRandom(); err != nil { + return err + } + } + return nil +} + +// MustSetRandom sets all elements of vector to random values, panicking if an error is encountered. +func (vector VectorE6) MustSetRandom() { + for i := range vector { + if _, err := vector[i].SetRandom(); err != nil { + panic(err) + } + } +} + // Butterfly computes the in-place butterfly operation on two vectors of E6 elements. // If other overlaps with vector, the result is undefined; the caller should use a temp vector. func (vector VectorE6) Butterfly(other VectorE6) { diff --git a/internal/generator/main.go b/internal/generator/main.go index de8d0ca99e..40e000da6e 100644 --- a/internal/generator/main.go +++ b/internal/generator/main.go @@ -84,14 +84,16 @@ func main() { field.WithIOP(), )) - // multilinear polynomials over the degree-6 extension E6 + // polynomial package (Polynomial, MultiLin, Pool, ...) over the + // degree-6 extension E6 if f.Name == "koalabear" { extInfo := fieldConfig.FieldDependency{ FieldPackagePath: "github.com/consensys/gnark-crypto/field/" + f.Name + "/extensions", FieldPackageName: "extensions", ElementType: "extensions.E6", + ExtensionDegree: 6, } - assertNoError(polynomial.GenerateMultilin(extInfo, filepath.Join(outputDir, "extensions", "polynomial"), true, gen)) + assertNoError(polynomial.Generate(extInfo, filepath.Join(outputDir, "extensions", "polynomial"), true, gen)) } }(conf) } diff --git a/internal/generator/polynomial/generate.go b/internal/generator/polynomial/generate.go index cfb0bdec70..7739d6b4ea 100644 --- a/internal/generator/polynomial/generate.go +++ b/internal/generator/polynomial/generate.go @@ -28,25 +28,3 @@ func Generate(conf config.FieldDependency, baseDir string, generateTests bool, g polyGen := common.NewDefaultGenerator(template.FS) return polyGen.Generate(conf, "polynomial", "", "", entries...) } - -// GenerateMultilin generates only the multilinear polynomial API (MultiLin and -// its memory Pool) over conf.ElementType. Unlike Generate, it does not require -// the element type to have the full base-field API (e.g. SetInt64, Vector -// helpers, big.Int conversions), so it can be used for field extensions such -// as koalabear's E6. -func GenerateMultilin(conf config.FieldDependency, baseDir string, generateTests bool, gen *common.Generator) error { - entries := []bavard.Entry{ - {File: filepath.Join(baseDir, "doc.go"), Templates: []string{"doc.multilin.go.tmpl"}}, - {File: filepath.Join(baseDir, "multilin.go"), Templates: []string{"multilin.go.tmpl"}}, - {File: filepath.Join(baseDir, "pool.go"), Templates: []string{"pool.go.tmpl"}}, - } - - if generateTests { - entries = append(entries, - bavard.Entry{File: filepath.Join(baseDir, "multilin_test.go"), Templates: []string{"multilin.test.go.tmpl"}}, - ) - } - - polyGen := common.NewDefaultGenerator(template.FS) - return polyGen.Generate(conf, "polynomial", "", "", entries...) -} diff --git a/internal/generator/polynomial/template/doc.multilin.go.tmpl b/internal/generator/polynomial/template/doc.multilin.go.tmpl deleted file mode 100644 index cf777371e9..0000000000 --- a/internal/generator/polynomial/template/doc.multilin.go.tmpl +++ /dev/null @@ -1,2 +0,0 @@ -// Package polynomial provides multilinear polynomial methods over {{.ElementType}}. -package polynomial diff --git a/internal/generator/polynomial/template/multilin.test.go.tmpl b/internal/generator/polynomial/template/multilin.test.go.tmpl index c9b0aa68b7..c71bf080f3 100644 --- a/internal/generator/polynomial/template/multilin.test.go.tmpl +++ b/internal/generator/polynomial/template/multilin.test.go.tmpl @@ -11,9 +11,11 @@ func TestFoldBilinear(t *testing.T) { // f = c₀ + c₁ X₁ + c₂ X₂ + c₃ X₁ X₂ var coefficients [4]{{.ElementType}} - for i := range coefficients { - coefficients[i].MustSetRandom() - } +{{- if .ExtensionDegree}} + {{.FieldPackageName}}.VectorE{{.ExtensionDegree}}(coefficients[:]).MustSetRandom() +{{- else}} + {{.FieldPackageName}}.Vector(coefficients[:]).MustSetRandom() +{{- end}} var r {{.ElementType}} r.MustSetRandom() @@ -49,12 +51,9 @@ func TestFoldBilinear(t *testing.T) { // TODO: Benchmark folding? Algorithms is pretty straightforward; unless we want to measure how well memory management is working func TestFoldedEqTable(t *testing.T) { - var one {{.ElementType}} - one.SetOne() - q := make([]{{.ElementType}}, 2) - q[0].Double(&one) // q₀ = 2 - q[1].Add(&q[0], &one) // q₁ = 3 + q[0].SetInt64(2) + q[1].SetInt64(3) m := make(MultiLin, 4) m[0].SetOne() @@ -63,6 +62,9 @@ func TestFoldedEqTable(t *testing.T) { eq := make([]{{.ElementType}}, 4) p := make([]{{.ElementType}}, 2) + var one {{.ElementType}} + one.SetOne() + for p0 := range 2 { p[1].SetZero() for p1 := range 2 { @@ -77,47 +79,3 @@ func TestFoldedEqTable(t *testing.T) { } } - -func TestEvaluateMatchesEqInnerProduct(t *testing.T) { - const n = 3 - - var one {{.ElementType}} - one.SetOne() - - m := make(MultiLin, 1<>(n-1-j))&1 == 1 { - b[j] = one - } else { - b[j].SetZero() - } - } - term = EvalEq(r, b) - term.Mul(&term, &m[i]) - expected.Add(&expected, &term) - } - - pool := NewPool(1 << n) - for _, p := range []*Pool{nil, &pool} { - clone := m.Clone() - res := m.Evaluate(r, p) - assert.True(t, res.Equal(&expected), "evaluation disagrees with Eq inner product") - for i := range m { - assert.True(t, m[i].Equal(&clone[i]), "Evaluate must not modify its receiver") - } - } - assert.Equal(t, n, m.NumVars()) -} diff --git a/internal/generator/polynomial/template/polynomial.go.tmpl b/internal/generator/polynomial/template/polynomial.go.tmpl index c0bec36089..0efdd422c6 100644 --- a/internal/generator/polynomial/template/polynomial.go.tmpl +++ b/internal/generator/polynomial/template/polynomial.go.tmpl @@ -148,6 +148,7 @@ func (p Polynomial) SetZero() { } } +{{if not .ExtensionDegree}} func (p Polynomial) Text(base int) string { var builder strings.Builder @@ -201,6 +202,7 @@ func (p Polynomial) Text(base int) string { return builder.String() } +{{end}} // InterpolateOnRange maps vector v to polynomial f // such that f(i) = v[i] for 0 ≤ i < len(v). @@ -293,7 +295,11 @@ func computeLagrangeBasis(domainSize uint8) []Polynomial { constTerms[l].Neg(&constTerms[l]) constTerms[l] = res[l].Eval(&constTerms[l]) } +{{- if .ExtensionDegree}} + constTerms = {{.FieldPackageName}}.BatchInvertE{{.ExtensionDegree}}(constTerms) +{{- else}} constTerms = {{.FieldPackageName}}.BatchInvert(constTerms) +{{- end}} for l := range domainSize { res[l].ScaleInPlace(&constTerms[l]) } diff --git a/internal/generator/polynomial/template/polynomial.test.go.tmpl b/internal/generator/polynomial/template/polynomial.test.go.tmpl index 1e82de85d3..da0af09257 100644 --- a/internal/generator/polynomial/template/polynomial.test.go.tmpl +++ b/internal/generator/polynomial/template/polynomial.test.go.tmpl @@ -192,6 +192,7 @@ func TestPolynomialAdd(t *testing.T) { } } +{{if not .ExtensionDegree}} func TestPolynomialText(t *testing.T) { var one, negTwo {{.ElementType}} one.SetOne() @@ -201,6 +202,7 @@ func TestPolynomialText(t *testing.T) { assert.Equal(t, "X² - 2X + 1", p.Text(10)) } +{{end}} func TestPrecomputeLagrange(t *testing.T) { @@ -214,7 +216,11 @@ func TestPrecomputeLagrange(t *testing.T) { y := polys[l].Eval(&I) if i == l && !y.IsOne() || i != l && !y.IsZero() { +{{- if .ExtensionDegree}} + t.Errorf("domainSize = %d: p_%d(%d) = %s", domainSize, l, i, y.String()) +{{- else}} t.Errorf("domainSize = %d: p_%d(%d) = %s", domainSize, l, i, y.Text(10)) +{{- end}} return false } } From bfa51d3c359fc01a91d640822295ee02021719ea Mon Sep 17 00:00:00 2001 From: Arya Tabaie Date: Thu, 1 Oct 2026 16:30:05 -0500 Subject: [PATCH 03/18] feat: `Bytes` for all extensions Signed-off-by: Arya Tabaie --- field/babybear/extensions/e2.go | 3 +++ field/babybear/extensions/e4.go | 3 +++ field/babybear/extensions/e6.go | 3 +++ field/babybear/extensions/vector.go | 20 +++++++++---------- field/goldilocks/extensions/e2.go | 3 +++ field/koalabear/extensions/e2.go | 3 +++ field/koalabear/extensions/e4.go | 3 +++ field/koalabear/extensions/e6.go | 3 +++ field/koalabear/extensions/vector.go | 20 +++++++++---------- field/mamabear/extensions/e3.go | 3 +++ .../field/template/extensions/e2.go.tmpl | 3 +++ .../field/template/extensions/e3.go.tmpl | 3 +++ .../field/template/extensions/e4.go.tmpl | 3 +++ .../field/template/extensions/e6.go.tmpl | 3 +++ .../field/template/extensions/vector.go.tmpl | 20 +++++++++---------- 15 files changed, 63 insertions(+), 33 deletions(-) diff --git a/field/babybear/extensions/e2.go b/field/babybear/extensions/e2.go index a143e7c6de..eed07c765a 100644 --- a/field/babybear/extensions/e2.go +++ b/field/babybear/extensions/e2.go @@ -11,6 +11,9 @@ import ( fr "github.com/consensys/gnark-crypto/field/babybear" ) +// BytesE2 is the number of bytes needed to represent a E2 +const BytesE2 = 2 * fr.Bytes + // E2 is a degree two finite field extension of fr.Element type E2 struct { A0, A1 fr.Element diff --git a/field/babybear/extensions/e4.go b/field/babybear/extensions/e4.go index a1e3495c6a..59fa646cd1 100644 --- a/field/babybear/extensions/e4.go +++ b/field/babybear/extensions/e4.go @@ -17,6 +17,9 @@ import ( const qInvNeg = 2013265919 const q = 2013265921 +// BytesE4 is the number of bytes needed to represent a E4 +const BytesE4 = 4 * fr.Bytes + // E4 is a degree two finite field extension of fr2 type E4 struct { B0, B1 E2 diff --git a/field/babybear/extensions/e6.go b/field/babybear/extensions/e6.go index afbfa69237..b57a1830ec 100644 --- a/field/babybear/extensions/e6.go +++ b/field/babybear/extensions/e6.go @@ -13,6 +13,9 @@ import ( fr "github.com/consensys/gnark-crypto/field/babybear" ) +// BytesE6 is the number of bytes needed to represent a E6 +const BytesE6 = 6 * fr.Bytes + // E6 is a degree three finite field extension of fp2 type E6 struct { B0, B1, B2 E2 diff --git a/field/babybear/extensions/vector.go b/field/babybear/extensions/vector.go index b42788f450..4991e9d2ab 100644 --- a/field/babybear/extensions/vector.go +++ b/field/babybear/extensions/vector.go @@ -238,11 +238,10 @@ func (vector *Vector) WriteTo(w io.Writer) (int64, error) { n := int64(4) - const e4Bytes = 4 * fr.Bytes - buf := make([]byte, len(*vector)*e4Bytes) + buf := make([]byte, len(*vector)*BytesE4) for i := range len(*vector) { - offset := i * e4Bytes + offset := i * BytesE4 fr.BigEndian.PutElement((*[fr.Bytes]byte)(buf[offset+0*fr.Bytes:offset+1*fr.Bytes]), (*vector)[i].B0.A0) fr.BigEndian.PutElement((*[fr.Bytes]byte)(buf[offset+1*fr.Bytes:offset+2*fr.Bytes]), (*vector)[i].B0.A1) fr.BigEndian.PutElement((*[fr.Bytes]byte)(buf[offset+2*fr.Bytes:offset+3*fr.Bytes]), (*vector)[i].B1.A0) @@ -290,9 +289,8 @@ func (vector *Vector) AsyncReadFrom(r io.Reader) (int64, error, chan error) { // return int64(read), err, chErr } headerSliceLen := uint64(binary.BigEndian.Uint32(bufSizeSlice[:])) - const e4Bytes = 4 * fr.Bytes if lr, ok := r.(interface{ Len() int }); ok { - if remaining := lr.Len(); remaining < 0 || headerSliceLen > uint64(remaining/e4Bytes) { + if remaining := lr.Len(); remaining < 0 || headerSliceLen > uint64(remaining/BytesE4) { close(chErr) return 4, io.ErrUnexpectedEOF, chErr } @@ -306,7 +304,7 @@ func (vector *Vector) AsyncReadFrom(r io.Reader) (int64, error, chan error) { // // reduce target size to 1GB on 32 bits architectures targetSize = uint64(1 << 30) // 1GB } - maxAllocateSliceLength := targetSize / uint64(e4Bytes) + maxAllocateSliceLength := targetSize / uint64(BytesE4) totalRead := int64(4) *vector = (*vector)[:0] @@ -322,12 +320,12 @@ func (vector *Vector) AsyncReadFrom(r io.Reader) (int64, error, chan error) { // if len(*vector) <= int(i) { *vector = append(*vector, make(Vector, int(min(headerSliceLen-i, maxAllocateSliceLength)))...) } - bSlice := unsafe.Slice((*byte)(unsafe.Pointer(&(*vector)[i])), int(min(headerSliceLen-i, maxAllocateSliceLength))*e4Bytes) + bSlice := unsafe.Slice((*byte)(unsafe.Pointer(&(*vector)[i])), int(min(headerSliceLen-i, maxAllocateSliceLength))*BytesE4) read, err := io.ReadFull(r, bSlice) totalRead += int64(read) if errors.Is(err, io.ErrUnexpectedEOF) { close(chErr) - return totalRead, fmt.Errorf("less data than expected: read %d elements, expected %d", i+uint64(read)/e4Bytes, headerSliceLen), chErr + return totalRead, fmt.Errorf("less data than expected: read %d elements, expected %d", i+uint64(read)/BytesE4, headerSliceLen), chErr } if err != nil { close(chErr) @@ -335,7 +333,7 @@ func (vector *Vector) AsyncReadFrom(r io.Reader) (int64, error, chan error) { // } } - bSlice := unsafe.Slice((*byte)(unsafe.Pointer(&(*vector)[0])), int(headerSliceLen)*e4Bytes) + bSlice := unsafe.Slice((*byte)(unsafe.Pointer(&(*vector)[0])), int(headerSliceLen)*BytesE4) go func() { setCoord := func(b *[fr.Bytes]byte) (fr.Element, bool) { e, err := fr.BigEndian.Element(b) @@ -349,8 +347,8 @@ func (vector *Vector) AsyncReadFrom(r io.Reader) (int64, error, chan error) { // var ok bool for i := range int(headerSliceLen) { - bstart := i * e4Bytes - bend := bstart + e4Bytes + bstart := i * BytesE4 + bend := bstart + BytesE4 b := bSlice[bstart:bend] (*vector)[i].B0.A0, ok = setCoord((*[fr.Bytes]byte)(b[0*fr.Bytes:])) diff --git a/field/goldilocks/extensions/e2.go b/field/goldilocks/extensions/e2.go index 95a38b8f64..3e81811e63 100644 --- a/field/goldilocks/extensions/e2.go +++ b/field/goldilocks/extensions/e2.go @@ -11,6 +11,9 @@ import ( fr "github.com/consensys/gnark-crypto/field/goldilocks" ) +// BytesE2 is the number of bytes needed to represent a E2 +const BytesE2 = 2 * fr.Bytes + // E2 is a degree two finite field extension of fr.Element type E2 struct { A0, A1 fr.Element diff --git a/field/koalabear/extensions/e2.go b/field/koalabear/extensions/e2.go index 7f3d442222..693e644120 100644 --- a/field/koalabear/extensions/e2.go +++ b/field/koalabear/extensions/e2.go @@ -11,6 +11,9 @@ import ( fr "github.com/consensys/gnark-crypto/field/koalabear" ) +// BytesE2 is the number of bytes needed to represent a E2 +const BytesE2 = 2 * fr.Bytes + // E2 is a degree two finite field extension of fr.Element type E2 struct { A0, A1 fr.Element diff --git a/field/koalabear/extensions/e4.go b/field/koalabear/extensions/e4.go index 6f7e8729a2..9d4be332af 100644 --- a/field/koalabear/extensions/e4.go +++ b/field/koalabear/extensions/e4.go @@ -17,6 +17,9 @@ import ( const qInvNeg = 2130706431 const q = 2130706433 +// BytesE4 is the number of bytes needed to represent a E4 +const BytesE4 = 4 * fr.Bytes + // E4 is a degree two finite field extension of fr2 type E4 struct { B0, B1 E2 diff --git a/field/koalabear/extensions/e6.go b/field/koalabear/extensions/e6.go index 964d524620..852d8073f4 100644 --- a/field/koalabear/extensions/e6.go +++ b/field/koalabear/extensions/e6.go @@ -14,6 +14,9 @@ import ( "github.com/consensys/gnark-crypto/utils/cpu" ) +// BytesE6 is the number of bytes needed to represent a E6 +const BytesE6 = 6 * fr.Bytes + // E6 is a degree three finite field extension of fp2 type E6 struct { B0, B1, B2 E2 diff --git a/field/koalabear/extensions/vector.go b/field/koalabear/extensions/vector.go index 6238c6d1ce..9f40b9a8d5 100644 --- a/field/koalabear/extensions/vector.go +++ b/field/koalabear/extensions/vector.go @@ -387,11 +387,10 @@ func (vector *Vector) WriteTo(w io.Writer) (int64, error) { n := int64(4) - const e4Bytes = 4 * fr.Bytes - buf := make([]byte, len(*vector)*e4Bytes) + buf := make([]byte, len(*vector)*BytesE4) for i := range len(*vector) { - offset := i * e4Bytes + offset := i * BytesE4 fr.BigEndian.PutElement((*[fr.Bytes]byte)(buf[offset+0*fr.Bytes:offset+1*fr.Bytes]), (*vector)[i].B0.A0) fr.BigEndian.PutElement((*[fr.Bytes]byte)(buf[offset+1*fr.Bytes:offset+2*fr.Bytes]), (*vector)[i].B0.A1) fr.BigEndian.PutElement((*[fr.Bytes]byte)(buf[offset+2*fr.Bytes:offset+3*fr.Bytes]), (*vector)[i].B1.A0) @@ -439,9 +438,8 @@ func (vector *Vector) AsyncReadFrom(r io.Reader) (int64, error, chan error) { // return int64(read), err, chErr } headerSliceLen := uint64(binary.BigEndian.Uint32(bufSizeSlice[:])) - const e4Bytes = 4 * fr.Bytes if lr, ok := r.(interface{ Len() int }); ok { - if remaining := lr.Len(); remaining < 0 || headerSliceLen > uint64(remaining/e4Bytes) { + if remaining := lr.Len(); remaining < 0 || headerSliceLen > uint64(remaining/BytesE4) { close(chErr) return 4, io.ErrUnexpectedEOF, chErr } @@ -455,7 +453,7 @@ func (vector *Vector) AsyncReadFrom(r io.Reader) (int64, error, chan error) { // // reduce target size to 1GB on 32 bits architectures targetSize = uint64(1 << 30) // 1GB } - maxAllocateSliceLength := targetSize / uint64(e4Bytes) + maxAllocateSliceLength := targetSize / uint64(BytesE4) totalRead := int64(4) *vector = (*vector)[:0] @@ -471,12 +469,12 @@ func (vector *Vector) AsyncReadFrom(r io.Reader) (int64, error, chan error) { // if len(*vector) <= int(i) { *vector = append(*vector, make(Vector, int(min(headerSliceLen-i, maxAllocateSliceLength)))...) } - bSlice := unsafe.Slice((*byte)(unsafe.Pointer(&(*vector)[i])), int(min(headerSliceLen-i, maxAllocateSliceLength))*e4Bytes) + bSlice := unsafe.Slice((*byte)(unsafe.Pointer(&(*vector)[i])), int(min(headerSliceLen-i, maxAllocateSliceLength))*BytesE4) read, err := io.ReadFull(r, bSlice) totalRead += int64(read) if errors.Is(err, io.ErrUnexpectedEOF) { close(chErr) - return totalRead, fmt.Errorf("less data than expected: read %d elements, expected %d", i+uint64(read)/e4Bytes, headerSliceLen), chErr + return totalRead, fmt.Errorf("less data than expected: read %d elements, expected %d", i+uint64(read)/BytesE4, headerSliceLen), chErr } if err != nil { close(chErr) @@ -484,7 +482,7 @@ func (vector *Vector) AsyncReadFrom(r io.Reader) (int64, error, chan error) { // } } - bSlice := unsafe.Slice((*byte)(unsafe.Pointer(&(*vector)[0])), int(headerSliceLen)*e4Bytes) + bSlice := unsafe.Slice((*byte)(unsafe.Pointer(&(*vector)[0])), int(headerSliceLen)*BytesE4) go func() { setCoord := func(b *[fr.Bytes]byte) (fr.Element, bool) { e, err := fr.BigEndian.Element(b) @@ -498,8 +496,8 @@ func (vector *Vector) AsyncReadFrom(r io.Reader) (int64, error, chan error) { // var ok bool for i := range int(headerSliceLen) { - bstart := i * e4Bytes - bend := bstart + e4Bytes + bstart := i * BytesE4 + bend := bstart + BytesE4 b := bSlice[bstart:bend] (*vector)[i].B0.A0, ok = setCoord((*[fr.Bytes]byte)(b[0*fr.Bytes:])) diff --git a/field/mamabear/extensions/e3.go b/field/mamabear/extensions/e3.go index 49305dd0b5..31aa8973e4 100644 --- a/field/mamabear/extensions/e3.go +++ b/field/mamabear/extensions/e3.go @@ -12,6 +12,9 @@ import ( fr "github.com/consensys/gnark-crypto/field/mamabear" ) +// BytesE3 is the number of bytes needed to represent a E3 +const BytesE3 = 3 * fr.Bytes + // E3 is an element of F_{p³} = F_p[t]/(t³ - t - 1). // Elements are represented as a0 + a1·t + a2·t² with a0, a1, a2 ∈ F_p. type E3 struct { diff --git a/internal/generator/field/template/extensions/e2.go.tmpl b/internal/generator/field/template/extensions/e2.go.tmpl index 9ed9fee4b3..9fb10db0a1 100644 --- a/internal/generator/field/template/extensions/e2.go.tmpl +++ b/internal/generator/field/template/extensions/e2.go.tmpl @@ -4,6 +4,9 @@ import ( fr "{{ .FieldPackagePath }}" ) +// BytesE2 is the number of bytes needed to represent a E2 +const BytesE2 = 2 * fr.Bytes + // E2 is a degree two finite field extension of fr.Element type E2 struct { A0, A1 fr.Element diff --git a/internal/generator/field/template/extensions/e3.go.tmpl b/internal/generator/field/template/extensions/e3.go.tmpl index a65441f03b..733bf7193a 100644 --- a/internal/generator/field/template/extensions/e3.go.tmpl +++ b/internal/generator/field/template/extensions/e3.go.tmpl @@ -5,6 +5,9 @@ import ( fr "{{ .FieldPackagePath }}" ) +// BytesE3 is the number of bytes needed to represent a E3 +const BytesE3 = 3 * fr.Bytes + // E3 is an element of F_{p³} = F_p[t]/(t³ - t - 1). // Elements are represented as a0 + a1·t + a2·t² with a0, a1, a2 ∈ F_p. type E3 struct { diff --git a/internal/generator/field/template/extensions/e4.go.tmpl b/internal/generator/field/template/extensions/e4.go.tmpl index 3f78c9e622..de9a7df975 100644 --- a/internal/generator/field/template/extensions/e4.go.tmpl +++ b/internal/generator/field/template/extensions/e4.go.tmpl @@ -10,6 +10,9 @@ import ( const qInvNeg = {{.QInvNeg}} const q = {{.Q}} +// BytesE4 is the number of bytes needed to represent a E4 +const BytesE4 = 4 * fr.Bytes + // E4 is a degree two finite field extension of fr2 type E4 struct { B0, B1 E2 diff --git a/internal/generator/field/template/extensions/e6.go.tmpl b/internal/generator/field/template/extensions/e6.go.tmpl index 0101c58129..a0149518f7 100644 --- a/internal/generator/field/template/extensions/e6.go.tmpl +++ b/internal/generator/field/template/extensions/e6.go.tmpl @@ -9,6 +9,9 @@ import ( {{- end }} ) +// BytesE6 is the number of bytes needed to represent a E6 +const BytesE6 = 6 * fr.Bytes + // E6 is a degree three finite field extension of fp2 type E6 struct { B0, B1, B2 E2 diff --git a/internal/generator/field/template/extensions/vector.go.tmpl b/internal/generator/field/template/extensions/vector.go.tmpl index c4b9c40005..1f6c7e8708 100644 --- a/internal/generator/field/template/extensions/vector.go.tmpl +++ b/internal/generator/field/template/extensions/vector.go.tmpl @@ -433,11 +433,10 @@ func (vector *Vector) WriteTo(w io.Writer) (int64, error) { n := int64(4) - const e4Bytes = 4 * fr.Bytes - buf := make([]byte, len(*vector)*e4Bytes) + buf := make([]byte, len(*vector)*BytesE4) for i := range len(*vector) { - offset := i * e4Bytes + offset := i * BytesE4 fr.BigEndian.PutElement((*[fr.Bytes]byte)(buf[offset+0*fr.Bytes:offset+1*fr.Bytes]), (*vector)[i].B0.A0) fr.BigEndian.PutElement((*[fr.Bytes]byte)(buf[offset+1*fr.Bytes:offset+2*fr.Bytes]), (*vector)[i].B0.A1) fr.BigEndian.PutElement((*[fr.Bytes]byte)(buf[offset+2*fr.Bytes:offset+3*fr.Bytes]), (*vector)[i].B1.A0) @@ -485,9 +484,8 @@ func (vector *Vector) AsyncReadFrom(r io.Reader) (int64, error, chan error) { // return int64(read), err, chErr } headerSliceLen := uint64(binary.BigEndian.Uint32(bufSizeSlice[:])) - const e4Bytes = 4 * fr.Bytes if lr, ok := r.(interface{ Len() int }); ok { - if remaining := lr.Len(); remaining < 0 || headerSliceLen > uint64(remaining/e4Bytes) { + if remaining := lr.Len(); remaining < 0 || headerSliceLen > uint64(remaining/BytesE4) { close(chErr) return 4, io.ErrUnexpectedEOF, chErr } @@ -501,7 +499,7 @@ func (vector *Vector) AsyncReadFrom(r io.Reader) (int64, error, chan error) { // // reduce target size to 1GB on 32 bits architectures targetSize = uint64(1 << 30) // 1GB } - maxAllocateSliceLength := targetSize / uint64(e4Bytes) + maxAllocateSliceLength := targetSize / uint64(BytesE4) totalRead := int64(4) *vector = (*vector)[:0] @@ -517,12 +515,12 @@ func (vector *Vector) AsyncReadFrom(r io.Reader) (int64, error, chan error) { // if len(*vector) <= int(i) { *vector = append(*vector, make(Vector, int(min(headerSliceLen-i, maxAllocateSliceLength)))...) } - bSlice := unsafe.Slice((*byte)(unsafe.Pointer(&(*vector)[i])), int(min(headerSliceLen-i, maxAllocateSliceLength))*e4Bytes) + bSlice := unsafe.Slice((*byte)(unsafe.Pointer(&(*vector)[i])), int(min(headerSliceLen-i, maxAllocateSliceLength))*BytesE4) read, err := io.ReadFull(r, bSlice) totalRead += int64(read) if errors.Is(err, io.ErrUnexpectedEOF) { close(chErr) - return totalRead, fmt.Errorf("less data than expected: read %d elements, expected %d", i+uint64(read)/e4Bytes, headerSliceLen), chErr + return totalRead, fmt.Errorf("less data than expected: read %d elements, expected %d", i+uint64(read)/BytesE4, headerSliceLen), chErr } if err != nil { close(chErr) @@ -530,7 +528,7 @@ func (vector *Vector) AsyncReadFrom(r io.Reader) (int64, error, chan error) { // } } - bSlice := unsafe.Slice((*byte)(unsafe.Pointer(&(*vector)[0])), int(headerSliceLen)*e4Bytes) + bSlice := unsafe.Slice((*byte)(unsafe.Pointer(&(*vector)[0])), int(headerSliceLen)*BytesE4) go func() { setCoord := func(b *[fr.Bytes]byte)(fr.Element, bool) { e, err := fr.BigEndian.Element(b) @@ -544,8 +542,8 @@ func (vector *Vector) AsyncReadFrom(r io.Reader) (int64, error, chan error) { // var ok bool for i := range int(headerSliceLen) { - bstart := i * e4Bytes - bend := bstart + e4Bytes + bstart := i * BytesE4 + bend := bstart + BytesE4 b := bSlice[bstart:bend] (*vector)[i].B0.A0, ok = setCoord((*[fr.Bytes]byte)(b[0*fr.Bytes:])) From 1a05ef28e702686f4caa5db857441992df13badf Mon Sep 17 00:00:00 2001 From: Arya Tabaie Date: Thu, 1 Oct 2026 22:06:45 -0500 Subject: [PATCH 04/18] feat: BigInt Signed-off-by: Arya Tabaie --- field/babybear/extensions/e2.go | 16 ++++++++++++++++ field/babybear/extensions/e4.go | 16 ++++++++++++++++ field/babybear/extensions/e6.go | 16 ++++++++++++++++ field/goldilocks/extensions/e2.go | 16 ++++++++++++++++ field/koalabear/extensions/e2.go | 16 ++++++++++++++++ field/koalabear/extensions/e4.go | 16 ++++++++++++++++ field/koalabear/extensions/e6.go | 16 ++++++++++++++++ field/mamabear/extensions/e3.go | 16 ++++++++++++++++ .../field/template/extensions/e2.go.tmpl | 16 ++++++++++++++++ .../field/template/extensions/e3.go.tmpl | 16 ++++++++++++++++ .../field/template/extensions/e4.go.tmpl | 16 ++++++++++++++++ .../field/template/extensions/e6.go.tmpl | 16 ++++++++++++++++ 12 files changed, 192 insertions(+) diff --git a/field/babybear/extensions/e2.go b/field/babybear/extensions/e2.go index eed07c765a..45811283bf 100644 --- a/field/babybear/extensions/e2.go +++ b/field/babybear/extensions/e2.go @@ -74,6 +74,22 @@ func (z *E2) SetOne() *E2 { return z } +// SetBigInt sets z to v (embedded via the unique ring homomorphism ℤ→𝔽, reduced modulo the base field order) and returns z +func (z *E2) SetBigInt(v *big.Int) *E2 { + *z = E2{} + z.A0.SetBigInt(v) + return z +} + +// BigInt sets res to the integer that z embeds, and returns res. +// It returns nil if z is not in the image of the embedding ℤ→𝔽, i.e. if any coordinate other than the first is non-zero. +func (z *E2) BigInt(res *big.Int) *big.Int { + if !(z.A1.IsZero()) { + return nil + } + return z.A0.BigInt(res) +} + // SetRandom sets a0 and a1 to random values func (z *E2) SetRandom() (*E2, error) { if _, err := z.A0.SetRandom(); err != nil { diff --git a/field/babybear/extensions/e4.go b/field/babybear/extensions/e4.go index 59fa646cd1..ca273926e7 100644 --- a/field/babybear/extensions/e4.go +++ b/field/babybear/extensions/e4.go @@ -85,6 +85,22 @@ func (z *E4) SetOne() *E4 { return z } +// SetBigInt sets z to v (embedded via the unique ring homomorphism ℤ→𝔽, reduced modulo the base field order) and returns z +func (z *E4) SetBigInt(v *big.Int) *E4 { + *z = E4{} + z.B0.A0.SetBigInt(v) + return z +} + +// BigInt sets res to the integer that z embeds, and returns res. +// It returns nil if z is not in the image of the embedding ℤ→𝔽, i.e. if any coordinate other than the first is non-zero. +func (z *E4) BigInt(res *big.Int) *big.Int { + if !(z.B0.A1.IsZero() && z.B1.A0.IsZero() && z.B1.A1.IsZero()) { + return nil + } + return z.B0.A0.BigInt(res) +} + // Lift sets the B0.A0 component of z to v func (z *E4) Lift(v *fr.Element) *E4 { *z = E4{} diff --git a/field/babybear/extensions/e6.go b/field/babybear/extensions/e6.go index b57a1830ec..78403df8c5 100644 --- a/field/babybear/extensions/e6.go +++ b/field/babybear/extensions/e6.go @@ -89,6 +89,22 @@ func (z *E6) SetUint64(v uint64) *E6 { return z } +// SetBigInt sets z to v (embedded via the unique ring homomorphism ℤ→𝔽, reduced modulo the base field order) and returns z +func (z *E6) SetBigInt(v *big.Int) *E6 { + *z = E6{} + z.B0.A0.SetBigInt(v) + return z +} + +// BigInt sets res to the integer that z embeds, and returns res. +// It returns nil if z is not in the image of the embedding ℤ→𝔽, i.e. if any coordinate other than the first is non-zero. +func (z *E6) BigInt(res *big.Int) *big.Int { + if !(z.B0.A1.IsZero() && z.B1.A0.IsZero() && z.B1.A1.IsZero() && z.B2.A0.IsZero() && z.B2.A1.IsZero()) { + return nil + } + return z.B0.A0.BigInt(res) +} + // MulByElement multiplies an element in E6 by an element in fr. // y may alias a coordinate of x, so we copy it first. func (z *E6) MulByElement(x *E6, y *fr.Element) *E6 { diff --git a/field/goldilocks/extensions/e2.go b/field/goldilocks/extensions/e2.go index 3e81811e63..0ed3ea21bb 100644 --- a/field/goldilocks/extensions/e2.go +++ b/field/goldilocks/extensions/e2.go @@ -74,6 +74,22 @@ func (z *E2) SetOne() *E2 { return z } +// SetBigInt sets z to v (embedded via the unique ring homomorphism ℤ→𝔽, reduced modulo the base field order) and returns z +func (z *E2) SetBigInt(v *big.Int) *E2 { + *z = E2{} + z.A0.SetBigInt(v) + return z +} + +// BigInt sets res to the integer that z embeds, and returns res. +// It returns nil if z is not in the image of the embedding ℤ→𝔽, i.e. if any coordinate other than the first is non-zero. +func (z *E2) BigInt(res *big.Int) *big.Int { + if !(z.A1.IsZero()) { + return nil + } + return z.A0.BigInt(res) +} + // SetRandom sets a0 and a1 to random values func (z *E2) SetRandom() (*E2, error) { if _, err := z.A0.SetRandom(); err != nil { diff --git a/field/koalabear/extensions/e2.go b/field/koalabear/extensions/e2.go index 693e644120..e757dd3b19 100644 --- a/field/koalabear/extensions/e2.go +++ b/field/koalabear/extensions/e2.go @@ -74,6 +74,22 @@ func (z *E2) SetOne() *E2 { return z } +// SetBigInt sets z to v (embedded via the unique ring homomorphism ℤ→𝔽, reduced modulo the base field order) and returns z +func (z *E2) SetBigInt(v *big.Int) *E2 { + *z = E2{} + z.A0.SetBigInt(v) + return z +} + +// BigInt sets res to the integer that z embeds, and returns res. +// It returns nil if z is not in the image of the embedding ℤ→𝔽, i.e. if any coordinate other than the first is non-zero. +func (z *E2) BigInt(res *big.Int) *big.Int { + if !(z.A1.IsZero()) { + return nil + } + return z.A0.BigInt(res) +} + // SetRandom sets a0 and a1 to random values func (z *E2) SetRandom() (*E2, error) { if _, err := z.A0.SetRandom(); err != nil { diff --git a/field/koalabear/extensions/e4.go b/field/koalabear/extensions/e4.go index 9d4be332af..ca335d07e7 100644 --- a/field/koalabear/extensions/e4.go +++ b/field/koalabear/extensions/e4.go @@ -85,6 +85,22 @@ func (z *E4) SetOne() *E4 { return z } +// SetBigInt sets z to v (embedded via the unique ring homomorphism ℤ→𝔽, reduced modulo the base field order) and returns z +func (z *E4) SetBigInt(v *big.Int) *E4 { + *z = E4{} + z.B0.A0.SetBigInt(v) + return z +} + +// BigInt sets res to the integer that z embeds, and returns res. +// It returns nil if z is not in the image of the embedding ℤ→𝔽, i.e. if any coordinate other than the first is non-zero. +func (z *E4) BigInt(res *big.Int) *big.Int { + if !(z.B0.A1.IsZero() && z.B1.A0.IsZero() && z.B1.A1.IsZero()) { + return nil + } + return z.B0.A0.BigInt(res) +} + // Lift sets the B0.A0 component of z to v func (z *E4) Lift(v *fr.Element) *E4 { *z = E4{} diff --git a/field/koalabear/extensions/e6.go b/field/koalabear/extensions/e6.go index 852d8073f4..8357c9ed6d 100644 --- a/field/koalabear/extensions/e6.go +++ b/field/koalabear/extensions/e6.go @@ -90,6 +90,22 @@ func (z *E6) SetUint64(v uint64) *E6 { return z } +// SetBigInt sets z to v (embedded via the unique ring homomorphism ℤ→𝔽, reduced modulo the base field order) and returns z +func (z *E6) SetBigInt(v *big.Int) *E6 { + *z = E6{} + z.B0.A0.SetBigInt(v) + return z +} + +// BigInt sets res to the integer that z embeds, and returns res. +// It returns nil if z is not in the image of the embedding ℤ→𝔽, i.e. if any coordinate other than the first is non-zero. +func (z *E6) BigInt(res *big.Int) *big.Int { + if !(z.B0.A1.IsZero() && z.B1.A0.IsZero() && z.B1.A1.IsZero() && z.B2.A0.IsZero() && z.B2.A1.IsZero()) { + return nil + } + return z.B0.A0.BigInt(res) +} + // MulByElement multiplies an element in E6 by an element in fr. // y may alias a coordinate of x, so we copy it first. func (z *E6) MulByElement(x *E6, y *fr.Element) *E6 { diff --git a/field/mamabear/extensions/e3.go b/field/mamabear/extensions/e3.go index 31aa8973e4..3459d04deb 100644 --- a/field/mamabear/extensions/e3.go +++ b/field/mamabear/extensions/e3.go @@ -42,6 +42,22 @@ func (z *E3) SetOne() *E3 { return z } +// SetBigInt sets z to v (embedded via the unique ring homomorphism ℤ→𝔽, reduced modulo the base field order) and returns z +func (z *E3) SetBigInt(v *big.Int) *E3 { + *z = E3{} + z.A0.SetBigInt(v) + return z +} + +// BigInt sets res to the integer that z embeds, and returns res. +// It returns nil if z is not in the image of the embedding ℤ→𝔽, i.e. if any coordinate other than the first is non-zero. +func (z *E3) BigInt(res *big.Int) *big.Int { + if !(z.A1.IsZero() && z.A2.IsZero()) { + return nil + } + return z.A0.BigInt(res) +} + // IsZero reports whether z is zero. func (z *E3) IsZero() bool { return z.A0.IsZero() && z.A1.IsZero() && z.A2.IsZero() diff --git a/internal/generator/field/template/extensions/e2.go.tmpl b/internal/generator/field/template/extensions/e2.go.tmpl index 9fb10db0a1..b2848b527f 100644 --- a/internal/generator/field/template/extensions/e2.go.tmpl +++ b/internal/generator/field/template/extensions/e2.go.tmpl @@ -67,6 +67,22 @@ func (z *E2) SetOne() *E2 { return z } +// SetBigInt sets z to v (embedded via the unique ring homomorphism ℤ→𝔽, reduced modulo the base field order) and returns z +func (z *E2) SetBigInt(v *big.Int) *E2 { + *z = E2{} + z.A0.SetBigInt(v) + return z +} + +// BigInt sets res to the integer that z embeds, and returns res. +// It returns nil if z is not in the image of the embedding ℤ→𝔽, i.e. if any coordinate other than the first is non-zero. +func (z *E2) BigInt(res *big.Int) *big.Int { + if !(z.A1.IsZero()) { + return nil + } + return z.A0.BigInt(res) +} + // SetRandom sets a0 and a1 to random values func (z *E2) SetRandom() (*E2, error) { if _, err := z.A0.SetRandom(); err != nil { diff --git a/internal/generator/field/template/extensions/e3.go.tmpl b/internal/generator/field/template/extensions/e3.go.tmpl index 733bf7193a..e521c1e50b 100644 --- a/internal/generator/field/template/extensions/e3.go.tmpl +++ b/internal/generator/field/template/extensions/e3.go.tmpl @@ -35,6 +35,22 @@ func (z *E3) SetOne() *E3 { return z } +// SetBigInt sets z to v (embedded via the unique ring homomorphism ℤ→𝔽, reduced modulo the base field order) and returns z +func (z *E3) SetBigInt(v *big.Int) *E3 { + *z = E3{} + z.A0.SetBigInt(v) + return z +} + +// BigInt sets res to the integer that z embeds, and returns res. +// It returns nil if z is not in the image of the embedding ℤ→𝔽, i.e. if any coordinate other than the first is non-zero. +func (z *E3) BigInt(res *big.Int) *big.Int { + if !(z.A1.IsZero() && z.A2.IsZero()) { + return nil + } + return z.A0.BigInt(res) +} + // IsZero reports whether z is zero. func (z *E3) IsZero() bool { return z.A0.IsZero() && z.A1.IsZero() && z.A2.IsZero() diff --git a/internal/generator/field/template/extensions/e4.go.tmpl b/internal/generator/field/template/extensions/e4.go.tmpl index de9a7df975..83e521cee1 100644 --- a/internal/generator/field/template/extensions/e4.go.tmpl +++ b/internal/generator/field/template/extensions/e4.go.tmpl @@ -78,6 +78,22 @@ func (z *E4) SetOne() *E4 { return z } +// SetBigInt sets z to v (embedded via the unique ring homomorphism ℤ→𝔽, reduced modulo the base field order) and returns z +func (z *E4) SetBigInt(v *big.Int) *E4 { + *z = E4{} + z.B0.A0.SetBigInt(v) + return z +} + +// BigInt sets res to the integer that z embeds, and returns res. +// It returns nil if z is not in the image of the embedding ℤ→𝔽, i.e. if any coordinate other than the first is non-zero. +func (z *E4) BigInt(res *big.Int) *big.Int { + if !(z.B0.A1.IsZero() && z.B1.A0.IsZero() && z.B1.A1.IsZero()) { + return nil + } + return z.B0.A0.BigInt(res) +} + // Lift sets the B0.A0 component of z to v func (z *E4) Lift(v *fr.Element) *E4 { *z = E4{} diff --git a/internal/generator/field/template/extensions/e6.go.tmpl b/internal/generator/field/template/extensions/e6.go.tmpl index a0149518f7..da23f8bcec 100644 --- a/internal/generator/field/template/extensions/e6.go.tmpl +++ b/internal/generator/field/template/extensions/e6.go.tmpl @@ -85,6 +85,22 @@ func (z *E6) SetUint64(v uint64) *E6 { return z } +// SetBigInt sets z to v (embedded via the unique ring homomorphism ℤ→𝔽, reduced modulo the base field order) and returns z +func (z *E6) SetBigInt(v *big.Int) *E6 { + *z = E6{} + z.B0.A0.SetBigInt(v) + return z +} + +// BigInt sets res to the integer that z embeds, and returns res. +// It returns nil if z is not in the image of the embedding ℤ→𝔽, i.e. if any coordinate other than the first is non-zero. +func (z *E6) BigInt(res *big.Int) *big.Int { + if !(z.B0.A1.IsZero() && z.B1.A0.IsZero() && z.B1.A1.IsZero() && z.B2.A0.IsZero() && z.B2.A1.IsZero()) { + return nil + } + return z.B0.A0.BigInt(res) +} + // MulByElement multiplies an element in E6 by an element in fr. // y may alias a coordinate of x, so we copy it first. func (z *E6) MulByElement(x *E6, y *fr.Element) *E6 { From 9e7ec3b0d8caee4d3710582df2d697f77fd45980 Mon Sep 17 00:00:00 2001 From: Arya Tabaie Date: Thu, 1 Oct 2026 22:11:04 -0500 Subject: [PATCH 05/18] feat: marshal and setbytes Signed-off-by: Arya Tabaie --- field/babybear/extensions/e6.go | 34 +++++++++++++++++ field/babybear/extensions/e6_test.go | 38 +++++++++++++++++++ field/koalabear/extensions/e6.go | 34 +++++++++++++++++ field/koalabear/extensions/e6_test.go | 38 +++++++++++++++++++ .../field/template/extensions/e6.go.tmpl | 34 +++++++++++++++++ .../field/template/extensions/e6_test.go.tmpl | 36 ++++++++++++++++++ 6 files changed, 214 insertions(+) diff --git a/field/babybear/extensions/e6.go b/field/babybear/extensions/e6.go index 78403df8c5..adfbbd4f73 100644 --- a/field/babybear/extensions/e6.go +++ b/field/babybear/extensions/e6.go @@ -105,6 +105,40 @@ func (z *E6) BigInt(res *big.Int) *big.Int { return z.B0.A0.BigInt(res) } +// Marshal returns the big-endian encodings of the coefficients +// B0.A0, B0.A1, B1.A0, B1.A1, B2.A0, B2.A1 concatenated, BytesE6 bytes in total +func (z *E6) Marshal() []byte { + res := make([]byte, 0, BytesE6) + b0 := z.B0.A0.Bytes() + res = append(res, b0[:]...) + b1 := z.B0.A1.Bytes() + res = append(res, b1[:]...) + b2 := z.B1.A0.Bytes() + res = append(res, b2[:]...) + b3 := z.B1.A1.Bytes() + res = append(res, b3[:]...) + b4 := z.B2.A0.Bytes() + res = append(res, b4[:]...) + b5 := z.B2.A1.Bytes() + res = append(res, b5[:]...) + return res +} + +// SetBytes sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytes, and returns z. +// It panics if len(b) != BytesE6. +func (z *E6) SetBytes(b []byte) *E6 { + if len(b) != BytesE6 { + panic("E6.SetBytes: invalid input length") + } + z.B0.A0.SetBytes(b[0*fr.Bytes : 1*fr.Bytes]) + z.B0.A1.SetBytes(b[1*fr.Bytes : 2*fr.Bytes]) + z.B1.A0.SetBytes(b[2*fr.Bytes : 3*fr.Bytes]) + z.B1.A1.SetBytes(b[3*fr.Bytes : 4*fr.Bytes]) + z.B2.A0.SetBytes(b[4*fr.Bytes : 5*fr.Bytes]) + z.B2.A1.SetBytes(b[5*fr.Bytes : 6*fr.Bytes]) + return z +} + // MulByElement multiplies an element in E6 by an element in fr. // y may alias a coordinate of x, so we copy it first. func (z *E6) MulByElement(x *E6, y *fr.Element) *E6 { diff --git a/field/babybear/extensions/e6_test.go b/field/babybear/extensions/e6_test.go index e323c413f9..51bf04fa8c 100644 --- a/field/babybear/extensions/e6_test.go +++ b/field/babybear/extensions/e6_test.go @@ -6,9 +6,11 @@ package extensions import ( + "bytes" "math/big" "testing" + fr "github.com/consensys/gnark-crypto/field/babybear" "github.com/leanovate/gopter" "github.com/leanovate/gopter/prop" ) @@ -357,3 +359,39 @@ func genE6() gopter.Gen { return E6{B0: values[0].(E2), B1: values[1].(E2), B2: values[2].(E2)} }) } + +func TestE6MarshalSetBytesRoundTrip(t *testing.T) { + for range 100 { + var x E6 + x.MustSetRandom() + + b := x.Marshal() + if len(b) != BytesE6 { + t.Fatalf("Marshal returned %d bytes, want %d", len(b), BytesE6) + } + coords := []*fr.Element{&x.B0.A0, &x.B0.A1, &x.B1.A0, &x.B1.A1, &x.B2.A0, &x.B2.A1} + for i, c := range coords { + cb := c.Bytes() + if !bytes.Equal(b[i*fr.Bytes:(i+1)*fr.Bytes], cb[:]) { + t.Fatalf("coefficient %d is not encoded at its expected offset", i) + } + } + + var y E6 + if !y.SetBytes(b).Equal(&x) { + t.Fatal("SetBytes(Marshal(x)) != x") + } + } + + for _, n := range []int{0, BytesE6 - 1, BytesE6 + 1} { + func() { + defer func() { + if recover() == nil { + t.Fatalf("SetBytes did not panic on %d bytes", n) + } + }() + var z E6 + z.SetBytes(make([]byte, n)) + }() + } +} diff --git a/field/koalabear/extensions/e6.go b/field/koalabear/extensions/e6.go index 8357c9ed6d..6e11ee4427 100644 --- a/field/koalabear/extensions/e6.go +++ b/field/koalabear/extensions/e6.go @@ -106,6 +106,40 @@ func (z *E6) BigInt(res *big.Int) *big.Int { return z.B0.A0.BigInt(res) } +// Marshal returns the big-endian encodings of the coefficients +// B0.A0, B0.A1, B1.A0, B1.A1, B2.A0, B2.A1 concatenated, BytesE6 bytes in total +func (z *E6) Marshal() []byte { + res := make([]byte, 0, BytesE6) + b0 := z.B0.A0.Bytes() + res = append(res, b0[:]...) + b1 := z.B0.A1.Bytes() + res = append(res, b1[:]...) + b2 := z.B1.A0.Bytes() + res = append(res, b2[:]...) + b3 := z.B1.A1.Bytes() + res = append(res, b3[:]...) + b4 := z.B2.A0.Bytes() + res = append(res, b4[:]...) + b5 := z.B2.A1.Bytes() + res = append(res, b5[:]...) + return res +} + +// SetBytes sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytes, and returns z. +// It panics if len(b) != BytesE6. +func (z *E6) SetBytes(b []byte) *E6 { + if len(b) != BytesE6 { + panic("E6.SetBytes: invalid input length") + } + z.B0.A0.SetBytes(b[0*fr.Bytes : 1*fr.Bytes]) + z.B0.A1.SetBytes(b[1*fr.Bytes : 2*fr.Bytes]) + z.B1.A0.SetBytes(b[2*fr.Bytes : 3*fr.Bytes]) + z.B1.A1.SetBytes(b[3*fr.Bytes : 4*fr.Bytes]) + z.B2.A0.SetBytes(b[4*fr.Bytes : 5*fr.Bytes]) + z.B2.A1.SetBytes(b[5*fr.Bytes : 6*fr.Bytes]) + return z +} + // MulByElement multiplies an element in E6 by an element in fr. // y may alias a coordinate of x, so we copy it first. func (z *E6) MulByElement(x *E6, y *fr.Element) *E6 { diff --git a/field/koalabear/extensions/e6_test.go b/field/koalabear/extensions/e6_test.go index 8f9745fd00..4fd52d1017 100644 --- a/field/koalabear/extensions/e6_test.go +++ b/field/koalabear/extensions/e6_test.go @@ -6,9 +6,11 @@ package extensions import ( + "bytes" "math/big" "testing" + fr "github.com/consensys/gnark-crypto/field/koalabear" "github.com/leanovate/gopter" "github.com/leanovate/gopter/prop" ) @@ -357,3 +359,39 @@ func genE6() gopter.Gen { return E6{B0: values[0].(E2), B1: values[1].(E2), B2: values[2].(E2)} }) } + +func TestE6MarshalSetBytesRoundTrip(t *testing.T) { + for range 100 { + var x E6 + x.MustSetRandom() + + b := x.Marshal() + if len(b) != BytesE6 { + t.Fatalf("Marshal returned %d bytes, want %d", len(b), BytesE6) + } + coords := []*fr.Element{&x.B0.A0, &x.B0.A1, &x.B1.A0, &x.B1.A1, &x.B2.A0, &x.B2.A1} + for i, c := range coords { + cb := c.Bytes() + if !bytes.Equal(b[i*fr.Bytes:(i+1)*fr.Bytes], cb[:]) { + t.Fatalf("coefficient %d is not encoded at its expected offset", i) + } + } + + var y E6 + if !y.SetBytes(b).Equal(&x) { + t.Fatal("SetBytes(Marshal(x)) != x") + } + } + + for _, n := range []int{0, BytesE6 - 1, BytesE6 + 1} { + func() { + defer func() { + if recover() == nil { + t.Fatalf("SetBytes did not panic on %d bytes", n) + } + }() + var z E6 + z.SetBytes(make([]byte, n)) + }() + } +} diff --git a/internal/generator/field/template/extensions/e6.go.tmpl b/internal/generator/field/template/extensions/e6.go.tmpl index da23f8bcec..4fd78f1820 100644 --- a/internal/generator/field/template/extensions/e6.go.tmpl +++ b/internal/generator/field/template/extensions/e6.go.tmpl @@ -101,6 +101,40 @@ func (z *E6) BigInt(res *big.Int) *big.Int { return z.B0.A0.BigInt(res) } +// Marshal returns the big-endian encodings of the coefficients +// B0.A0, B0.A1, B1.A0, B1.A1, B2.A0, B2.A1 concatenated, BytesE6 bytes in total +func (z *E6) Marshal() []byte { + res := make([]byte, 0, BytesE6) + b0 := z.B0.A0.Bytes() + res = append(res, b0[:]...) + b1 := z.B0.A1.Bytes() + res = append(res, b1[:]...) + b2 := z.B1.A0.Bytes() + res = append(res, b2[:]...) + b3 := z.B1.A1.Bytes() + res = append(res, b3[:]...) + b4 := z.B2.A0.Bytes() + res = append(res, b4[:]...) + b5 := z.B2.A1.Bytes() + res = append(res, b5[:]...) + return res +} + +// SetBytes sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytes, and returns z. +// It panics if len(b) != BytesE6. +func (z *E6) SetBytes(b []byte) *E6 { + if len(b) != BytesE6 { + panic("E6.SetBytes: invalid input length") + } + z.B0.A0.SetBytes(b[0*fr.Bytes : 1*fr.Bytes]) + z.B0.A1.SetBytes(b[1*fr.Bytes : 2*fr.Bytes]) + z.B1.A0.SetBytes(b[2*fr.Bytes : 3*fr.Bytes]) + z.B1.A1.SetBytes(b[3*fr.Bytes : 4*fr.Bytes]) + z.B2.A0.SetBytes(b[4*fr.Bytes : 5*fr.Bytes]) + z.B2.A1.SetBytes(b[5*fr.Bytes : 6*fr.Bytes]) + return z +} + // MulByElement multiplies an element in E6 by an element in fr. // y may alias a coordinate of x, so we copy it first. func (z *E6) MulByElement(x *E6, y *fr.Element) *E6 { diff --git a/internal/generator/field/template/extensions/e6_test.go.tmpl b/internal/generator/field/template/extensions/e6_test.go.tmpl index ae7ab28eb0..9de29ff567 100644 --- a/internal/generator/field/template/extensions/e6_test.go.tmpl +++ b/internal/generator/field/template/extensions/e6_test.go.tmpl @@ -350,3 +350,39 @@ func genE6() gopter.Gen { return E6{B0: values[0].(E2), B1: values[1].(E2), B2: values[2].(E2)} }) } + +func TestE6MarshalSetBytesRoundTrip(t *testing.T) { + for range 100 { + var x E6 + x.MustSetRandom() + + b := x.Marshal() + if len(b) != BytesE6 { + t.Fatalf("Marshal returned %d bytes, want %d", len(b), BytesE6) + } + coords := []*fr.Element{&x.B0.A0, &x.B0.A1, &x.B1.A0, &x.B1.A1, &x.B2.A0, &x.B2.A1} + for i, c := range coords { + cb := c.Bytes() + if !bytes.Equal(b[i*fr.Bytes:(i+1)*fr.Bytes], cb[:]) { + t.Fatalf("coefficient %d is not encoded at its expected offset", i) + } + } + + var y E6 + if !y.SetBytes(b).Equal(&x) { + t.Fatal("SetBytes(Marshal(x)) != x") + } + } + + for _, n := range []int{0, BytesE6 - 1, BytesE6 + 1} { + func() { + defer func() { + if recover() == nil { + t.Fatalf("SetBytes did not panic on %d bytes", n) + } + }() + var z E6 + z.SetBytes(make([]byte, n)) + }() + } +} From e9438bf24ae985f631ea6c66deab6c349daad317 Mon Sep 17 00:00:00 2001 From: Arya Tabaie Date: Thu, 1 Oct 2026 22:16:53 -0500 Subject: [PATCH 06/18] feat: marshal and setbytes for all extensions Signed-off-by: Arya Tabaie --- field/babybear/extensions/e2.go | 22 ++++++++++- field/babybear/extensions/e2_test.go | 37 +++++++++++++++++++ field/babybear/extensions/e4.go | 26 ++++++++++++- field/babybear/extensions/e4_test.go | 36 ++++++++++++++++++ field/babybear/extensions/e6.go | 22 ++++------- field/goldilocks/extensions/e2.go | 22 ++++++++++- field/goldilocks/extensions/e2_test.go | 37 +++++++++++++++++++ field/koalabear/extensions/e2.go | 22 ++++++++++- field/koalabear/extensions/e2_test.go | 37 +++++++++++++++++++ field/koalabear/extensions/e4.go | 26 ++++++++++++- field/koalabear/extensions/e4_test.go | 36 ++++++++++++++++++ field/koalabear/extensions/e6.go | 22 ++++------- field/mamabear/extensions/e3.go | 24 +++++++++++- field/mamabear/extensions/e3_test.go | 37 +++++++++++++++++++ .../field/template/extensions/e2.go.tmpl | 22 ++++++++++- .../field/template/extensions/e2_test.go.tmpl | 36 ++++++++++++++++++ .../field/template/extensions/e3.go.tmpl | 24 +++++++++++- .../field/template/extensions/e3_test.go.tmpl | 36 ++++++++++++++++++ .../field/template/extensions/e4.go.tmpl | 26 ++++++++++++- .../field/template/extensions/e4_test.go.tmpl | 36 ++++++++++++++++++ .../field/template/extensions/e6.go.tmpl | 22 ++++------- 21 files changed, 557 insertions(+), 51 deletions(-) diff --git a/field/babybear/extensions/e2.go b/field/babybear/extensions/e2.go index 45811283bf..51648d6d46 100644 --- a/field/babybear/extensions/e2.go +++ b/field/babybear/extensions/e2.go @@ -84,12 +84,32 @@ func (z *E2) SetBigInt(v *big.Int) *E2 { // BigInt sets res to the integer that z embeds, and returns res. // It returns nil if z is not in the image of the embedding ℤ→𝔽, i.e. if any coordinate other than the first is non-zero. func (z *E2) BigInt(res *big.Int) *big.Int { - if !(z.A1.IsZero()) { + if !z.A1.IsZero() { return nil } return z.A0.BigInt(res) } +// Marshal returns the big-endian encodings of the coefficients +// A0, A1 concatenated, BytesE2 bytes in total +func (z *E2) Marshal() []byte { + res := make([]byte, BytesE2) + fr.BigEndian.PutElement((*[fr.Bytes]byte)(res[0*fr.Bytes:1*fr.Bytes]), z.A0) + fr.BigEndian.PutElement((*[fr.Bytes]byte)(res[1*fr.Bytes:2*fr.Bytes]), z.A1) + return res +} + +// SetBytes sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytes, and returns z. +// It panics if len(b) != BytesE2. +func (z *E2) SetBytes(b []byte) *E2 { + if len(b) != BytesE2 { + panic("E2.SetBytes: invalid input length") + } + z.A0.SetBytes(b[0*fr.Bytes : 1*fr.Bytes]) + z.A1.SetBytes(b[1*fr.Bytes : 2*fr.Bytes]) + return z +} + // SetRandom sets a0 and a1 to random values func (z *E2) SetRandom() (*E2, error) { if _, err := z.A0.SetRandom(); err != nil { diff --git a/field/babybear/extensions/e2_test.go b/field/babybear/extensions/e2_test.go index 82091e834e..7bbc971c76 100644 --- a/field/babybear/extensions/e2_test.go +++ b/field/babybear/extensions/e2_test.go @@ -6,6 +6,7 @@ package extensions import ( + "bytes" "crypto/rand" "math/big" "testing" @@ -553,3 +554,39 @@ func genE2() gopter.Gen { return E2{A0: values[0].(fr.Element), A1: values[1].(fr.Element)} }) } + +func TestE2MarshalSetBytesRoundTrip(t *testing.T) { + for range 100 { + var x E2 + x.MustSetRandom() + + b := x.Marshal() + if len(b) != BytesE2 { + t.Fatalf("Marshal returned %d bytes, want %d", len(b), BytesE2) + } + coords := []*fr.Element{&x.A0, &x.A1} + for i, c := range coords { + cb := c.Bytes() + if !bytes.Equal(b[i*fr.Bytes:(i+1)*fr.Bytes], cb[:]) { + t.Fatalf("coefficient %d is not encoded at its expected offset", i) + } + } + + var y E2 + if !y.SetBytes(b).Equal(&x) { + t.Fatal("SetBytes(Marshal(x)) != x") + } + } + + for _, n := range []int{0, BytesE2 - 1, BytesE2 + 1} { + func() { + defer func() { + if recover() == nil { + t.Fatalf("SetBytes did not panic on %d bytes", n) + } + }() + var z E2 + z.SetBytes(make([]byte, n)) + }() + } +} diff --git a/field/babybear/extensions/e4.go b/field/babybear/extensions/e4.go index ca273926e7..349ff59e82 100644 --- a/field/babybear/extensions/e4.go +++ b/field/babybear/extensions/e4.go @@ -95,12 +95,36 @@ func (z *E4) SetBigInt(v *big.Int) *E4 { // BigInt sets res to the integer that z embeds, and returns res. // It returns nil if z is not in the image of the embedding ℤ→𝔽, i.e. if any coordinate other than the first is non-zero. func (z *E4) BigInt(res *big.Int) *big.Int { - if !(z.B0.A1.IsZero() && z.B1.A0.IsZero() && z.B1.A1.IsZero()) { + if !z.B0.A1.IsZero() || !z.B1.A0.IsZero() || !z.B1.A1.IsZero() { return nil } return z.B0.A0.BigInt(res) } +// Marshal returns the big-endian encodings of the coefficients +// B0.A0, B0.A1, B1.A0, B1.A1 concatenated, BytesE4 bytes in total +func (z *E4) Marshal() []byte { + res := make([]byte, BytesE4) + fr.BigEndian.PutElement((*[fr.Bytes]byte)(res[0*fr.Bytes:1*fr.Bytes]), z.B0.A0) + fr.BigEndian.PutElement((*[fr.Bytes]byte)(res[1*fr.Bytes:2*fr.Bytes]), z.B0.A1) + fr.BigEndian.PutElement((*[fr.Bytes]byte)(res[2*fr.Bytes:3*fr.Bytes]), z.B1.A0) + fr.BigEndian.PutElement((*[fr.Bytes]byte)(res[3*fr.Bytes:4*fr.Bytes]), z.B1.A1) + return res +} + +// SetBytes sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytes, and returns z. +// It panics if len(b) != BytesE4. +func (z *E4) SetBytes(b []byte) *E4 { + if len(b) != BytesE4 { + panic("E4.SetBytes: invalid input length") + } + z.B0.A0.SetBytes(b[0*fr.Bytes : 1*fr.Bytes]) + z.B0.A1.SetBytes(b[1*fr.Bytes : 2*fr.Bytes]) + z.B1.A0.SetBytes(b[2*fr.Bytes : 3*fr.Bytes]) + z.B1.A1.SetBytes(b[3*fr.Bytes : 4*fr.Bytes]) + return z +} + // Lift sets the B0.A0 component of z to v func (z *E4) Lift(v *fr.Element) *E4 { *z = E4{} diff --git a/field/babybear/extensions/e4_test.go b/field/babybear/extensions/e4_test.go index 94abaccda9..bc3d43d252 100644 --- a/field/babybear/extensions/e4_test.go +++ b/field/babybear/extensions/e4_test.go @@ -1126,3 +1126,39 @@ func genFrVector(size int) gopter.Gen { return gopter.NewGenResult(v, gopter.NoShrinker) } } + +func TestE4MarshalSetBytesRoundTrip(t *testing.T) { + for range 100 { + var x E4 + x.MustSetRandom() + + b := x.Marshal() + if len(b) != BytesE4 { + t.Fatalf("Marshal returned %d bytes, want %d", len(b), BytesE4) + } + coords := []*fr.Element{&x.B0.A0, &x.B0.A1, &x.B1.A0, &x.B1.A1} + for i, c := range coords { + cb := c.Bytes() + if !bytes.Equal(b[i*fr.Bytes:(i+1)*fr.Bytes], cb[:]) { + t.Fatalf("coefficient %d is not encoded at its expected offset", i) + } + } + + var y E4 + if !y.SetBytes(b).Equal(&x) { + t.Fatal("SetBytes(Marshal(x)) != x") + } + } + + for _, n := range []int{0, BytesE4 - 1, BytesE4 + 1} { + func() { + defer func() { + if recover() == nil { + t.Fatalf("SetBytes did not panic on %d bytes", n) + } + }() + var z E4 + z.SetBytes(make([]byte, n)) + }() + } +} diff --git a/field/babybear/extensions/e6.go b/field/babybear/extensions/e6.go index adfbbd4f73..c4af4e1847 100644 --- a/field/babybear/extensions/e6.go +++ b/field/babybear/extensions/e6.go @@ -99,7 +99,7 @@ func (z *E6) SetBigInt(v *big.Int) *E6 { // BigInt sets res to the integer that z embeds, and returns res. // It returns nil if z is not in the image of the embedding ℤ→𝔽, i.e. if any coordinate other than the first is non-zero. func (z *E6) BigInt(res *big.Int) *big.Int { - if !(z.B0.A1.IsZero() && z.B1.A0.IsZero() && z.B1.A1.IsZero() && z.B2.A0.IsZero() && z.B2.A1.IsZero()) { + if !z.B0.A1.IsZero() || !z.B1.A0.IsZero() || !z.B1.A1.IsZero() || !z.B2.A0.IsZero() || !z.B2.A1.IsZero() { return nil } return z.B0.A0.BigInt(res) @@ -108,19 +108,13 @@ func (z *E6) BigInt(res *big.Int) *big.Int { // Marshal returns the big-endian encodings of the coefficients // B0.A0, B0.A1, B1.A0, B1.A1, B2.A0, B2.A1 concatenated, BytesE6 bytes in total func (z *E6) Marshal() []byte { - res := make([]byte, 0, BytesE6) - b0 := z.B0.A0.Bytes() - res = append(res, b0[:]...) - b1 := z.B0.A1.Bytes() - res = append(res, b1[:]...) - b2 := z.B1.A0.Bytes() - res = append(res, b2[:]...) - b3 := z.B1.A1.Bytes() - res = append(res, b3[:]...) - b4 := z.B2.A0.Bytes() - res = append(res, b4[:]...) - b5 := z.B2.A1.Bytes() - res = append(res, b5[:]...) + res := make([]byte, BytesE6) + fr.BigEndian.PutElement((*[fr.Bytes]byte)(res[0*fr.Bytes:1*fr.Bytes]), z.B0.A0) + fr.BigEndian.PutElement((*[fr.Bytes]byte)(res[1*fr.Bytes:2*fr.Bytes]), z.B0.A1) + fr.BigEndian.PutElement((*[fr.Bytes]byte)(res[2*fr.Bytes:3*fr.Bytes]), z.B1.A0) + fr.BigEndian.PutElement((*[fr.Bytes]byte)(res[3*fr.Bytes:4*fr.Bytes]), z.B1.A1) + fr.BigEndian.PutElement((*[fr.Bytes]byte)(res[4*fr.Bytes:5*fr.Bytes]), z.B2.A0) + fr.BigEndian.PutElement((*[fr.Bytes]byte)(res[5*fr.Bytes:6*fr.Bytes]), z.B2.A1) return res } diff --git a/field/goldilocks/extensions/e2.go b/field/goldilocks/extensions/e2.go index 0ed3ea21bb..d673316f42 100644 --- a/field/goldilocks/extensions/e2.go +++ b/field/goldilocks/extensions/e2.go @@ -84,12 +84,32 @@ func (z *E2) SetBigInt(v *big.Int) *E2 { // BigInt sets res to the integer that z embeds, and returns res. // It returns nil if z is not in the image of the embedding ℤ→𝔽, i.e. if any coordinate other than the first is non-zero. func (z *E2) BigInt(res *big.Int) *big.Int { - if !(z.A1.IsZero()) { + if !z.A1.IsZero() { return nil } return z.A0.BigInt(res) } +// Marshal returns the big-endian encodings of the coefficients +// A0, A1 concatenated, BytesE2 bytes in total +func (z *E2) Marshal() []byte { + res := make([]byte, BytesE2) + fr.BigEndian.PutElement((*[fr.Bytes]byte)(res[0*fr.Bytes:1*fr.Bytes]), z.A0) + fr.BigEndian.PutElement((*[fr.Bytes]byte)(res[1*fr.Bytes:2*fr.Bytes]), z.A1) + return res +} + +// SetBytes sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytes, and returns z. +// It panics if len(b) != BytesE2. +func (z *E2) SetBytes(b []byte) *E2 { + if len(b) != BytesE2 { + panic("E2.SetBytes: invalid input length") + } + z.A0.SetBytes(b[0*fr.Bytes : 1*fr.Bytes]) + z.A1.SetBytes(b[1*fr.Bytes : 2*fr.Bytes]) + return z +} + // SetRandom sets a0 and a1 to random values func (z *E2) SetRandom() (*E2, error) { if _, err := z.A0.SetRandom(); err != nil { diff --git a/field/goldilocks/extensions/e2_test.go b/field/goldilocks/extensions/e2_test.go index 418a2cd785..0cac5a0f66 100644 --- a/field/goldilocks/extensions/e2_test.go +++ b/field/goldilocks/extensions/e2_test.go @@ -6,6 +6,7 @@ package extensions import ( + "bytes" "crypto/rand" "math/big" "testing" @@ -536,3 +537,39 @@ func genE2() gopter.Gen { return E2{A0: values[0].(fr.Element), A1: values[1].(fr.Element)} }) } + +func TestE2MarshalSetBytesRoundTrip(t *testing.T) { + for range 100 { + var x E2 + x.MustSetRandom() + + b := x.Marshal() + if len(b) != BytesE2 { + t.Fatalf("Marshal returned %d bytes, want %d", len(b), BytesE2) + } + coords := []*fr.Element{&x.A0, &x.A1} + for i, c := range coords { + cb := c.Bytes() + if !bytes.Equal(b[i*fr.Bytes:(i+1)*fr.Bytes], cb[:]) { + t.Fatalf("coefficient %d is not encoded at its expected offset", i) + } + } + + var y E2 + if !y.SetBytes(b).Equal(&x) { + t.Fatal("SetBytes(Marshal(x)) != x") + } + } + + for _, n := range []int{0, BytesE2 - 1, BytesE2 + 1} { + func() { + defer func() { + if recover() == nil { + t.Fatalf("SetBytes did not panic on %d bytes", n) + } + }() + var z E2 + z.SetBytes(make([]byte, n)) + }() + } +} diff --git a/field/koalabear/extensions/e2.go b/field/koalabear/extensions/e2.go index e757dd3b19..e2f27d899c 100644 --- a/field/koalabear/extensions/e2.go +++ b/field/koalabear/extensions/e2.go @@ -84,12 +84,32 @@ func (z *E2) SetBigInt(v *big.Int) *E2 { // BigInt sets res to the integer that z embeds, and returns res. // It returns nil if z is not in the image of the embedding ℤ→𝔽, i.e. if any coordinate other than the first is non-zero. func (z *E2) BigInt(res *big.Int) *big.Int { - if !(z.A1.IsZero()) { + if !z.A1.IsZero() { return nil } return z.A0.BigInt(res) } +// Marshal returns the big-endian encodings of the coefficients +// A0, A1 concatenated, BytesE2 bytes in total +func (z *E2) Marshal() []byte { + res := make([]byte, BytesE2) + fr.BigEndian.PutElement((*[fr.Bytes]byte)(res[0*fr.Bytes:1*fr.Bytes]), z.A0) + fr.BigEndian.PutElement((*[fr.Bytes]byte)(res[1*fr.Bytes:2*fr.Bytes]), z.A1) + return res +} + +// SetBytes sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytes, and returns z. +// It panics if len(b) != BytesE2. +func (z *E2) SetBytes(b []byte) *E2 { + if len(b) != BytesE2 { + panic("E2.SetBytes: invalid input length") + } + z.A0.SetBytes(b[0*fr.Bytes : 1*fr.Bytes]) + z.A1.SetBytes(b[1*fr.Bytes : 2*fr.Bytes]) + return z +} + // SetRandom sets a0 and a1 to random values func (z *E2) SetRandom() (*E2, error) { if _, err := z.A0.SetRandom(); err != nil { diff --git a/field/koalabear/extensions/e2_test.go b/field/koalabear/extensions/e2_test.go index 6d69d55e4a..048a52a055 100644 --- a/field/koalabear/extensions/e2_test.go +++ b/field/koalabear/extensions/e2_test.go @@ -6,6 +6,7 @@ package extensions import ( + "bytes" "crypto/rand" "math/big" "testing" @@ -553,3 +554,39 @@ func genE2() gopter.Gen { return E2{A0: values[0].(fr.Element), A1: values[1].(fr.Element)} }) } + +func TestE2MarshalSetBytesRoundTrip(t *testing.T) { + for range 100 { + var x E2 + x.MustSetRandom() + + b := x.Marshal() + if len(b) != BytesE2 { + t.Fatalf("Marshal returned %d bytes, want %d", len(b), BytesE2) + } + coords := []*fr.Element{&x.A0, &x.A1} + for i, c := range coords { + cb := c.Bytes() + if !bytes.Equal(b[i*fr.Bytes:(i+1)*fr.Bytes], cb[:]) { + t.Fatalf("coefficient %d is not encoded at its expected offset", i) + } + } + + var y E2 + if !y.SetBytes(b).Equal(&x) { + t.Fatal("SetBytes(Marshal(x)) != x") + } + } + + for _, n := range []int{0, BytesE2 - 1, BytesE2 + 1} { + func() { + defer func() { + if recover() == nil { + t.Fatalf("SetBytes did not panic on %d bytes", n) + } + }() + var z E2 + z.SetBytes(make([]byte, n)) + }() + } +} diff --git a/field/koalabear/extensions/e4.go b/field/koalabear/extensions/e4.go index ca335d07e7..3622804b3a 100644 --- a/field/koalabear/extensions/e4.go +++ b/field/koalabear/extensions/e4.go @@ -95,12 +95,36 @@ func (z *E4) SetBigInt(v *big.Int) *E4 { // BigInt sets res to the integer that z embeds, and returns res. // It returns nil if z is not in the image of the embedding ℤ→𝔽, i.e. if any coordinate other than the first is non-zero. func (z *E4) BigInt(res *big.Int) *big.Int { - if !(z.B0.A1.IsZero() && z.B1.A0.IsZero() && z.B1.A1.IsZero()) { + if !z.B0.A1.IsZero() || !z.B1.A0.IsZero() || !z.B1.A1.IsZero() { return nil } return z.B0.A0.BigInt(res) } +// Marshal returns the big-endian encodings of the coefficients +// B0.A0, B0.A1, B1.A0, B1.A1 concatenated, BytesE4 bytes in total +func (z *E4) Marshal() []byte { + res := make([]byte, BytesE4) + fr.BigEndian.PutElement((*[fr.Bytes]byte)(res[0*fr.Bytes:1*fr.Bytes]), z.B0.A0) + fr.BigEndian.PutElement((*[fr.Bytes]byte)(res[1*fr.Bytes:2*fr.Bytes]), z.B0.A1) + fr.BigEndian.PutElement((*[fr.Bytes]byte)(res[2*fr.Bytes:3*fr.Bytes]), z.B1.A0) + fr.BigEndian.PutElement((*[fr.Bytes]byte)(res[3*fr.Bytes:4*fr.Bytes]), z.B1.A1) + return res +} + +// SetBytes sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytes, and returns z. +// It panics if len(b) != BytesE4. +func (z *E4) SetBytes(b []byte) *E4 { + if len(b) != BytesE4 { + panic("E4.SetBytes: invalid input length") + } + z.B0.A0.SetBytes(b[0*fr.Bytes : 1*fr.Bytes]) + z.B0.A1.SetBytes(b[1*fr.Bytes : 2*fr.Bytes]) + z.B1.A0.SetBytes(b[2*fr.Bytes : 3*fr.Bytes]) + z.B1.A1.SetBytes(b[3*fr.Bytes : 4*fr.Bytes]) + return z +} + // Lift sets the B0.A0 component of z to v func (z *E4) Lift(v *fr.Element) *E4 { *z = E4{} diff --git a/field/koalabear/extensions/e4_test.go b/field/koalabear/extensions/e4_test.go index ee883f7302..6394b3d676 100644 --- a/field/koalabear/extensions/e4_test.go +++ b/field/koalabear/extensions/e4_test.go @@ -1126,3 +1126,39 @@ func genFrVector(size int) gopter.Gen { return gopter.NewGenResult(v, gopter.NoShrinker) } } + +func TestE4MarshalSetBytesRoundTrip(t *testing.T) { + for range 100 { + var x E4 + x.MustSetRandom() + + b := x.Marshal() + if len(b) != BytesE4 { + t.Fatalf("Marshal returned %d bytes, want %d", len(b), BytesE4) + } + coords := []*fr.Element{&x.B0.A0, &x.B0.A1, &x.B1.A0, &x.B1.A1} + for i, c := range coords { + cb := c.Bytes() + if !bytes.Equal(b[i*fr.Bytes:(i+1)*fr.Bytes], cb[:]) { + t.Fatalf("coefficient %d is not encoded at its expected offset", i) + } + } + + var y E4 + if !y.SetBytes(b).Equal(&x) { + t.Fatal("SetBytes(Marshal(x)) != x") + } + } + + for _, n := range []int{0, BytesE4 - 1, BytesE4 + 1} { + func() { + defer func() { + if recover() == nil { + t.Fatalf("SetBytes did not panic on %d bytes", n) + } + }() + var z E4 + z.SetBytes(make([]byte, n)) + }() + } +} diff --git a/field/koalabear/extensions/e6.go b/field/koalabear/extensions/e6.go index 6e11ee4427..ad7f664ed9 100644 --- a/field/koalabear/extensions/e6.go +++ b/field/koalabear/extensions/e6.go @@ -100,7 +100,7 @@ func (z *E6) SetBigInt(v *big.Int) *E6 { // BigInt sets res to the integer that z embeds, and returns res. // It returns nil if z is not in the image of the embedding ℤ→𝔽, i.e. if any coordinate other than the first is non-zero. func (z *E6) BigInt(res *big.Int) *big.Int { - if !(z.B0.A1.IsZero() && z.B1.A0.IsZero() && z.B1.A1.IsZero() && z.B2.A0.IsZero() && z.B2.A1.IsZero()) { + if !z.B0.A1.IsZero() || !z.B1.A0.IsZero() || !z.B1.A1.IsZero() || !z.B2.A0.IsZero() || !z.B2.A1.IsZero() { return nil } return z.B0.A0.BigInt(res) @@ -109,19 +109,13 @@ func (z *E6) BigInt(res *big.Int) *big.Int { // Marshal returns the big-endian encodings of the coefficients // B0.A0, B0.A1, B1.A0, B1.A1, B2.A0, B2.A1 concatenated, BytesE6 bytes in total func (z *E6) Marshal() []byte { - res := make([]byte, 0, BytesE6) - b0 := z.B0.A0.Bytes() - res = append(res, b0[:]...) - b1 := z.B0.A1.Bytes() - res = append(res, b1[:]...) - b2 := z.B1.A0.Bytes() - res = append(res, b2[:]...) - b3 := z.B1.A1.Bytes() - res = append(res, b3[:]...) - b4 := z.B2.A0.Bytes() - res = append(res, b4[:]...) - b5 := z.B2.A1.Bytes() - res = append(res, b5[:]...) + res := make([]byte, BytesE6) + fr.BigEndian.PutElement((*[fr.Bytes]byte)(res[0*fr.Bytes:1*fr.Bytes]), z.B0.A0) + fr.BigEndian.PutElement((*[fr.Bytes]byte)(res[1*fr.Bytes:2*fr.Bytes]), z.B0.A1) + fr.BigEndian.PutElement((*[fr.Bytes]byte)(res[2*fr.Bytes:3*fr.Bytes]), z.B1.A0) + fr.BigEndian.PutElement((*[fr.Bytes]byte)(res[3*fr.Bytes:4*fr.Bytes]), z.B1.A1) + fr.BigEndian.PutElement((*[fr.Bytes]byte)(res[4*fr.Bytes:5*fr.Bytes]), z.B2.A0) + fr.BigEndian.PutElement((*[fr.Bytes]byte)(res[5*fr.Bytes:6*fr.Bytes]), z.B2.A1) return res } diff --git a/field/mamabear/extensions/e3.go b/field/mamabear/extensions/e3.go index 3459d04deb..84fd32b301 100644 --- a/field/mamabear/extensions/e3.go +++ b/field/mamabear/extensions/e3.go @@ -52,12 +52,34 @@ func (z *E3) SetBigInt(v *big.Int) *E3 { // BigInt sets res to the integer that z embeds, and returns res. // It returns nil if z is not in the image of the embedding ℤ→𝔽, i.e. if any coordinate other than the first is non-zero. func (z *E3) BigInt(res *big.Int) *big.Int { - if !(z.A1.IsZero() && z.A2.IsZero()) { + if !z.A1.IsZero() || !z.A2.IsZero() { return nil } return z.A0.BigInt(res) } +// Marshal returns the big-endian encodings of the coefficients +// A0, A1, A2 concatenated, BytesE3 bytes in total +func (z *E3) Marshal() []byte { + res := make([]byte, BytesE3) + fr.BigEndian.PutElement((*[fr.Bytes]byte)(res[0*fr.Bytes:1*fr.Bytes]), z.A0) + fr.BigEndian.PutElement((*[fr.Bytes]byte)(res[1*fr.Bytes:2*fr.Bytes]), z.A1) + fr.BigEndian.PutElement((*[fr.Bytes]byte)(res[2*fr.Bytes:3*fr.Bytes]), z.A2) + return res +} + +// SetBytes sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytes, and returns z. +// It panics if len(b) != BytesE3. +func (z *E3) SetBytes(b []byte) *E3 { + if len(b) != BytesE3 { + panic("E3.SetBytes: invalid input length") + } + z.A0.SetBytes(b[0*fr.Bytes : 1*fr.Bytes]) + z.A1.SetBytes(b[1*fr.Bytes : 2*fr.Bytes]) + z.A2.SetBytes(b[2*fr.Bytes : 3*fr.Bytes]) + return z +} + // IsZero reports whether z is zero. func (z *E3) IsZero() bool { return z.A0.IsZero() && z.A1.IsZero() && z.A2.IsZero() diff --git a/field/mamabear/extensions/e3_test.go b/field/mamabear/extensions/e3_test.go index 2b370b66e1..bd8fe05104 100644 --- a/field/mamabear/extensions/e3_test.go +++ b/field/mamabear/extensions/e3_test.go @@ -6,6 +6,7 @@ package extensions import ( + "bytes" "math/big" "testing" @@ -442,3 +443,39 @@ func BenchmarkE3VectorMulAccByElement(b *testing.B) { dst.MulAccByElement(scale, &alpha) } } + +func TestE3MarshalSetBytesRoundTrip(t *testing.T) { + for range 100 { + var x E3 + x.MustSetRandom() + + b := x.Marshal() + if len(b) != BytesE3 { + t.Fatalf("Marshal returned %d bytes, want %d", len(b), BytesE3) + } + coords := []*fr.Element{&x.A0, &x.A1, &x.A2} + for i, c := range coords { + cb := c.Bytes() + if !bytes.Equal(b[i*fr.Bytes:(i+1)*fr.Bytes], cb[:]) { + t.Fatalf("coefficient %d is not encoded at its expected offset", i) + } + } + + var y E3 + if !y.SetBytes(b).Equal(&x) { + t.Fatal("SetBytes(Marshal(x)) != x") + } + } + + for _, n := range []int{0, BytesE3 - 1, BytesE3 + 1} { + func() { + defer func() { + if recover() == nil { + t.Fatalf("SetBytes did not panic on %d bytes", n) + } + }() + var z E3 + z.SetBytes(make([]byte, n)) + }() + } +} diff --git a/internal/generator/field/template/extensions/e2.go.tmpl b/internal/generator/field/template/extensions/e2.go.tmpl index b2848b527f..8d5daf2bf6 100644 --- a/internal/generator/field/template/extensions/e2.go.tmpl +++ b/internal/generator/field/template/extensions/e2.go.tmpl @@ -77,12 +77,32 @@ func (z *E2) SetBigInt(v *big.Int) *E2 { // BigInt sets res to the integer that z embeds, and returns res. // It returns nil if z is not in the image of the embedding ℤ→𝔽, i.e. if any coordinate other than the first is non-zero. func (z *E2) BigInt(res *big.Int) *big.Int { - if !(z.A1.IsZero()) { + if !z.A1.IsZero() { return nil } return z.A0.BigInt(res) } +// Marshal returns the big-endian encodings of the coefficients +// A0, A1 concatenated, BytesE2 bytes in total +func (z *E2) Marshal() []byte { + res := make([]byte, BytesE2) + fr.BigEndian.PutElement((*[fr.Bytes]byte)(res[0*fr.Bytes:1*fr.Bytes]), z.A0) + fr.BigEndian.PutElement((*[fr.Bytes]byte)(res[1*fr.Bytes:2*fr.Bytes]), z.A1) + return res +} + +// SetBytes sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytes, and returns z. +// It panics if len(b) != BytesE2. +func (z *E2) SetBytes(b []byte) *E2 { + if len(b) != BytesE2 { + panic("E2.SetBytes: invalid input length") + } + z.A0.SetBytes(b[0*fr.Bytes : 1*fr.Bytes]) + z.A1.SetBytes(b[1*fr.Bytes : 2*fr.Bytes]) + return z +} + // SetRandom sets a0 and a1 to random values func (z *E2) SetRandom() (*E2, error) { if _, err := z.A0.SetRandom(); err != nil { diff --git a/internal/generator/field/template/extensions/e2_test.go.tmpl b/internal/generator/field/template/extensions/e2_test.go.tmpl index 4c75f673eb..e01d8e2b8a 100644 --- a/internal/generator/field/template/extensions/e2_test.go.tmpl +++ b/internal/generator/field/template/extensions/e2_test.go.tmpl @@ -559,3 +559,39 @@ func genE2() gopter.Gen { return E2{A0: values[0].(fr.Element), A1: values[1].(fr.Element)} }) } + +func TestE2MarshalSetBytesRoundTrip(t *testing.T) { + for range 100 { + var x E2 + x.MustSetRandom() + + b := x.Marshal() + if len(b) != BytesE2 { + t.Fatalf("Marshal returned %d bytes, want %d", len(b), BytesE2) + } + coords := []*fr.Element{&x.A0, &x.A1} + for i, c := range coords { + cb := c.Bytes() + if !bytes.Equal(b[i*fr.Bytes:(i+1)*fr.Bytes], cb[:]) { + t.Fatalf("coefficient %d is not encoded at its expected offset", i) + } + } + + var y E2 + if !y.SetBytes(b).Equal(&x) { + t.Fatal("SetBytes(Marshal(x)) != x") + } + } + + for _, n := range []int{0, BytesE2 - 1, BytesE2 + 1} { + func() { + defer func() { + if recover() == nil { + t.Fatalf("SetBytes did not panic on %d bytes", n) + } + }() + var z E2 + z.SetBytes(make([]byte, n)) + }() + } +} diff --git a/internal/generator/field/template/extensions/e3.go.tmpl b/internal/generator/field/template/extensions/e3.go.tmpl index e521c1e50b..c6fcb64913 100644 --- a/internal/generator/field/template/extensions/e3.go.tmpl +++ b/internal/generator/field/template/extensions/e3.go.tmpl @@ -45,12 +45,34 @@ func (z *E3) SetBigInt(v *big.Int) *E3 { // BigInt sets res to the integer that z embeds, and returns res. // It returns nil if z is not in the image of the embedding ℤ→𝔽, i.e. if any coordinate other than the first is non-zero. func (z *E3) BigInt(res *big.Int) *big.Int { - if !(z.A1.IsZero() && z.A2.IsZero()) { + if !z.A1.IsZero() || !z.A2.IsZero() { return nil } return z.A0.BigInt(res) } +// Marshal returns the big-endian encodings of the coefficients +// A0, A1, A2 concatenated, BytesE3 bytes in total +func (z *E3) Marshal() []byte { + res := make([]byte, BytesE3) + fr.BigEndian.PutElement((*[fr.Bytes]byte)(res[0*fr.Bytes:1*fr.Bytes]), z.A0) + fr.BigEndian.PutElement((*[fr.Bytes]byte)(res[1*fr.Bytes:2*fr.Bytes]), z.A1) + fr.BigEndian.PutElement((*[fr.Bytes]byte)(res[2*fr.Bytes:3*fr.Bytes]), z.A2) + return res +} + +// SetBytes sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytes, and returns z. +// It panics if len(b) != BytesE3. +func (z *E3) SetBytes(b []byte) *E3 { + if len(b) != BytesE3 { + panic("E3.SetBytes: invalid input length") + } + z.A0.SetBytes(b[0*fr.Bytes : 1*fr.Bytes]) + z.A1.SetBytes(b[1*fr.Bytes : 2*fr.Bytes]) + z.A2.SetBytes(b[2*fr.Bytes : 3*fr.Bytes]) + return z +} + // IsZero reports whether z is zero. func (z *E3) IsZero() bool { return z.A0.IsZero() && z.A1.IsZero() && z.A2.IsZero() diff --git a/internal/generator/field/template/extensions/e3_test.go.tmpl b/internal/generator/field/template/extensions/e3_test.go.tmpl index 8f85e316fc..18a93c0625 100644 --- a/internal/generator/field/template/extensions/e3_test.go.tmpl +++ b/internal/generator/field/template/extensions/e3_test.go.tmpl @@ -435,3 +435,39 @@ func BenchmarkE3VectorMulAccByElement(b *testing.B) { dst.MulAccByElement(scale, &alpha) } } + +func TestE3MarshalSetBytesRoundTrip(t *testing.T) { + for range 100 { + var x E3 + x.MustSetRandom() + + b := x.Marshal() + if len(b) != BytesE3 { + t.Fatalf("Marshal returned %d bytes, want %d", len(b), BytesE3) + } + coords := []*fr.Element{&x.A0, &x.A1, &x.A2} + for i, c := range coords { + cb := c.Bytes() + if !bytes.Equal(b[i*fr.Bytes:(i+1)*fr.Bytes], cb[:]) { + t.Fatalf("coefficient %d is not encoded at its expected offset", i) + } + } + + var y E3 + if !y.SetBytes(b).Equal(&x) { + t.Fatal("SetBytes(Marshal(x)) != x") + } + } + + for _, n := range []int{0, BytesE3 - 1, BytesE3 + 1} { + func() { + defer func() { + if recover() == nil { + t.Fatalf("SetBytes did not panic on %d bytes", n) + } + }() + var z E3 + z.SetBytes(make([]byte, n)) + }() + } +} diff --git a/internal/generator/field/template/extensions/e4.go.tmpl b/internal/generator/field/template/extensions/e4.go.tmpl index 83e521cee1..3c6ba0b4b1 100644 --- a/internal/generator/field/template/extensions/e4.go.tmpl +++ b/internal/generator/field/template/extensions/e4.go.tmpl @@ -88,12 +88,36 @@ func (z *E4) SetBigInt(v *big.Int) *E4 { // BigInt sets res to the integer that z embeds, and returns res. // It returns nil if z is not in the image of the embedding ℤ→𝔽, i.e. if any coordinate other than the first is non-zero. func (z *E4) BigInt(res *big.Int) *big.Int { - if !(z.B0.A1.IsZero() && z.B1.A0.IsZero() && z.B1.A1.IsZero()) { + if !z.B0.A1.IsZero() || !z.B1.A0.IsZero() || !z.B1.A1.IsZero() { return nil } return z.B0.A0.BigInt(res) } +// Marshal returns the big-endian encodings of the coefficients +// B0.A0, B0.A1, B1.A0, B1.A1 concatenated, BytesE4 bytes in total +func (z *E4) Marshal() []byte { + res := make([]byte, BytesE4) + fr.BigEndian.PutElement((*[fr.Bytes]byte)(res[0*fr.Bytes:1*fr.Bytes]), z.B0.A0) + fr.BigEndian.PutElement((*[fr.Bytes]byte)(res[1*fr.Bytes:2*fr.Bytes]), z.B0.A1) + fr.BigEndian.PutElement((*[fr.Bytes]byte)(res[2*fr.Bytes:3*fr.Bytes]), z.B1.A0) + fr.BigEndian.PutElement((*[fr.Bytes]byte)(res[3*fr.Bytes:4*fr.Bytes]), z.B1.A1) + return res +} + +// SetBytes sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytes, and returns z. +// It panics if len(b) != BytesE4. +func (z *E4) SetBytes(b []byte) *E4 { + if len(b) != BytesE4 { + panic("E4.SetBytes: invalid input length") + } + z.B0.A0.SetBytes(b[0*fr.Bytes : 1*fr.Bytes]) + z.B0.A1.SetBytes(b[1*fr.Bytes : 2*fr.Bytes]) + z.B1.A0.SetBytes(b[2*fr.Bytes : 3*fr.Bytes]) + z.B1.A1.SetBytes(b[3*fr.Bytes : 4*fr.Bytes]) + return z +} + // Lift sets the B0.A0 component of z to v func (z *E4) Lift(v *fr.Element) *E4 { *z = E4{} diff --git a/internal/generator/field/template/extensions/e4_test.go.tmpl b/internal/generator/field/template/extensions/e4_test.go.tmpl index 2500313a20..ac2879da37 100644 --- a/internal/generator/field/template/extensions/e4_test.go.tmpl +++ b/internal/generator/field/template/extensions/e4_test.go.tmpl @@ -1124,3 +1124,39 @@ func genFrVector(size int) gopter.Gen { return gopter.NewGenResult(v, gopter.NoShrinker) } } + +func TestE4MarshalSetBytesRoundTrip(t *testing.T) { + for range 100 { + var x E4 + x.MustSetRandom() + + b := x.Marshal() + if len(b) != BytesE4 { + t.Fatalf("Marshal returned %d bytes, want %d", len(b), BytesE4) + } + coords := []*fr.Element{&x.B0.A0, &x.B0.A1, &x.B1.A0, &x.B1.A1} + for i, c := range coords { + cb := c.Bytes() + if !bytes.Equal(b[i*fr.Bytes:(i+1)*fr.Bytes], cb[:]) { + t.Fatalf("coefficient %d is not encoded at its expected offset", i) + } + } + + var y E4 + if !y.SetBytes(b).Equal(&x) { + t.Fatal("SetBytes(Marshal(x)) != x") + } + } + + for _, n := range []int{0, BytesE4 - 1, BytesE4 + 1} { + func() { + defer func() { + if recover() == nil { + t.Fatalf("SetBytes did not panic on %d bytes", n) + } + }() + var z E4 + z.SetBytes(make([]byte, n)) + }() + } +} diff --git a/internal/generator/field/template/extensions/e6.go.tmpl b/internal/generator/field/template/extensions/e6.go.tmpl index 4fd78f1820..f3008945f2 100644 --- a/internal/generator/field/template/extensions/e6.go.tmpl +++ b/internal/generator/field/template/extensions/e6.go.tmpl @@ -95,7 +95,7 @@ func (z *E6) SetBigInt(v *big.Int) *E6 { // BigInt sets res to the integer that z embeds, and returns res. // It returns nil if z is not in the image of the embedding ℤ→𝔽, i.e. if any coordinate other than the first is non-zero. func (z *E6) BigInt(res *big.Int) *big.Int { - if !(z.B0.A1.IsZero() && z.B1.A0.IsZero() && z.B1.A1.IsZero() && z.B2.A0.IsZero() && z.B2.A1.IsZero()) { + if !z.B0.A1.IsZero() || !z.B1.A0.IsZero() || !z.B1.A1.IsZero() || !z.B2.A0.IsZero() || !z.B2.A1.IsZero() { return nil } return z.B0.A0.BigInt(res) @@ -104,19 +104,13 @@ func (z *E6) BigInt(res *big.Int) *big.Int { // Marshal returns the big-endian encodings of the coefficients // B0.A0, B0.A1, B1.A0, B1.A1, B2.A0, B2.A1 concatenated, BytesE6 bytes in total func (z *E6) Marshal() []byte { - res := make([]byte, 0, BytesE6) - b0 := z.B0.A0.Bytes() - res = append(res, b0[:]...) - b1 := z.B0.A1.Bytes() - res = append(res, b1[:]...) - b2 := z.B1.A0.Bytes() - res = append(res, b2[:]...) - b3 := z.B1.A1.Bytes() - res = append(res, b3[:]...) - b4 := z.B2.A0.Bytes() - res = append(res, b4[:]...) - b5 := z.B2.A1.Bytes() - res = append(res, b5[:]...) + res := make([]byte, BytesE6) + fr.BigEndian.PutElement((*[fr.Bytes]byte)(res[0*fr.Bytes:1*fr.Bytes]), z.B0.A0) + fr.BigEndian.PutElement((*[fr.Bytes]byte)(res[1*fr.Bytes:2*fr.Bytes]), z.B0.A1) + fr.BigEndian.PutElement((*[fr.Bytes]byte)(res[2*fr.Bytes:3*fr.Bytes]), z.B1.A0) + fr.BigEndian.PutElement((*[fr.Bytes]byte)(res[3*fr.Bytes:4*fr.Bytes]), z.B1.A1) + fr.BigEndian.PutElement((*[fr.Bytes]byte)(res[4*fr.Bytes:5*fr.Bytes]), z.B2.A0) + fr.BigEndian.PutElement((*[fr.Bytes]byte)(res[5*fr.Bytes:6*fr.Bytes]), z.B2.A1) return res } From 454ab6773cb2fea7e00026bfabd02e671420ee6b Mon Sep 17 00:00:00 2001 From: Arya Tabaie Date: Sun, 4 Oct 2026 21:23:52 -0500 Subject: [PATCH 07/18] feat: PR feedback Signed-off-by: Arya Tabaie --- field/babybear/extensions/e2.go | 9 ++- field/babybear/extensions/e2_test.go | 19 +++-- field/babybear/extensions/e4.go | 9 ++- field/babybear/extensions/e4_test.go | 17 ++--- field/babybear/extensions/e6.go | 9 ++- field/babybear/extensions/e6_test.go | 19 +++-- field/goldilocks/extensions/e2.go | 9 ++- field/goldilocks/extensions/e2_test.go | 19 +++-- field/koalabear/extensions/e2.go | 9 ++- field/koalabear/extensions/e2_test.go | 19 +++-- field/koalabear/extensions/e4.go | 9 ++- field/koalabear/extensions/e4_test.go | 17 ++--- field/koalabear/extensions/e6.go | 9 ++- field/koalabear/extensions/e6_test.go | 19 +++-- .../{multilin.go => multilin_e6.go} | 34 ++++----- .../{multilin_test.go => multilin_e6_test.go} | 10 +-- .../{polynomial.go => polynomial_e6.go} | 72 +++++++++--------- ...lynomial_test.go => polynomial_e6_test.go} | 40 +++++----- .../polynomial/{pool.go => pool_e6.go} | 70 ++++++++--------- field/mamabear/extensions/e3.go | 9 ++- field/mamabear/extensions/e3_test.go | 19 +++-- internal/generator/config/fields.go | 9 ++- .../generator/field/config/field_config.go | 10 +++ .../field/template/extensions/e2.go.tmpl | 8 +- .../field/template/extensions/e2_test.go.tmpl | 18 ++--- .../field/template/extensions/e3.go.tmpl | 8 +- .../field/template/extensions/e3_test.go.tmpl | 18 ++--- .../field/template/extensions/e4.go.tmpl | 8 +- .../field/template/extensions/e4_test.go.tmpl | 17 ++--- .../field/template/extensions/e6.go.tmpl | 8 +- .../field/template/extensions/e6_test.go.tmpl | 18 ++--- internal/generator/main.go | 12 +-- internal/generator/polynomial/generate.go | 31 ++++++-- .../polynomial/template/multilin.go.tmpl | 34 ++++----- .../polynomial/template/multilin.test.go.tmpl | 12 +-- .../polynomial/template/polynomial.go.tmpl | 76 +++++++++---------- .../template/polynomial.test.go.tmpl | 44 +++++------ .../polynomial/template/pool.go.tmpl | 70 ++++++++--------- 38 files changed, 432 insertions(+), 415 deletions(-) rename field/koalabear/extensions/polynomial/{multilin.go => multilin_e6.go} (85%) rename field/koalabear/extensions/polynomial/{multilin_test.go => multilin_e6_test.go} (91%) rename field/koalabear/extensions/polynomial/{polynomial.go => polynomial_e6.go} (71%) rename field/koalabear/extensions/polynomial/{polynomial_test.go => polynomial_e6_test.go} (84%) rename field/koalabear/extensions/polynomial/{pool.go => pool_e6.go} (67%) diff --git a/field/babybear/extensions/e2.go b/field/babybear/extensions/e2.go index 51648d6d46..033a3e8ef0 100644 --- a/field/babybear/extensions/e2.go +++ b/field/babybear/extensions/e2.go @@ -6,6 +6,7 @@ package extensions import ( + "fmt" "math/big" fr "github.com/consensys/gnark-crypto/field/babybear" @@ -100,14 +101,14 @@ func (z *E2) Marshal() []byte { } // SetBytes sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytes, and returns z. -// It panics if len(b) != BytesE2. -func (z *E2) SetBytes(b []byte) *E2 { +// It returns an error if len(b) != BytesE2. +func (z *E2) SetBytes(b []byte) (*E2, error) { if len(b) != BytesE2 { - panic("E2.SetBytes: invalid input length") + return nil, fmt.Errorf("E2.SetBytes: got %d bytes, expected %d", len(b), BytesE2) } z.A0.SetBytes(b[0*fr.Bytes : 1*fr.Bytes]) z.A1.SetBytes(b[1*fr.Bytes : 2*fr.Bytes]) - return z + return z, nil } // SetRandom sets a0 and a1 to random values diff --git a/field/babybear/extensions/e2_test.go b/field/babybear/extensions/e2_test.go index 7bbc971c76..05e44f1ee6 100644 --- a/field/babybear/extensions/e2_test.go +++ b/field/babybear/extensions/e2_test.go @@ -11,6 +11,8 @@ import ( "math/big" "testing" + "github.com/stretchr/testify/require" + fr "github.com/consensys/gnark-crypto/field/babybear" "github.com/leanovate/gopter" "github.com/leanovate/gopter/prop" @@ -573,20 +575,17 @@ func TestE2MarshalSetBytesRoundTrip(t *testing.T) { } var y E2 - if !y.SetBytes(b).Equal(&x) { + _, err := y.SetBytes(b) + require.NoError(t, err) + if !y.Equal(&x) { t.Fatal("SetBytes(Marshal(x)) != x") } } for _, n := range []int{0, BytesE2 - 1, BytesE2 + 1} { - func() { - defer func() { - if recover() == nil { - t.Fatalf("SetBytes did not panic on %d bytes", n) - } - }() - var z E2 - z.SetBytes(make([]byte, n)) - }() + var z E2 + if _, err := z.SetBytes(make([]byte, n)); err == nil { + t.Fatalf("SetBytes did not fail on %d bytes", n) + } } } diff --git a/field/babybear/extensions/e4.go b/field/babybear/extensions/e4.go index 349ff59e82..f6aba77ab2 100644 --- a/field/babybear/extensions/e4.go +++ b/field/babybear/extensions/e4.go @@ -6,6 +6,7 @@ package extensions import ( + "fmt" "math/big" "math/bits" @@ -113,16 +114,16 @@ func (z *E4) Marshal() []byte { } // SetBytes sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytes, and returns z. -// It panics if len(b) != BytesE4. -func (z *E4) SetBytes(b []byte) *E4 { +// It returns an error if len(b) != BytesE4. +func (z *E4) SetBytes(b []byte) (*E4, error) { if len(b) != BytesE4 { - panic("E4.SetBytes: invalid input length") + return nil, fmt.Errorf("E4.SetBytes: got %d bytes, expected %d", len(b), BytesE4) } z.B0.A0.SetBytes(b[0*fr.Bytes : 1*fr.Bytes]) z.B0.A1.SetBytes(b[1*fr.Bytes : 2*fr.Bytes]) z.B1.A0.SetBytes(b[2*fr.Bytes : 3*fr.Bytes]) z.B1.A1.SetBytes(b[3*fr.Bytes : 4*fr.Bytes]) - return z + return z, nil } // Lift sets the B0.A0 component of z to v diff --git a/field/babybear/extensions/e4_test.go b/field/babybear/extensions/e4_test.go index bc3d43d252..4becccf8d7 100644 --- a/field/babybear/extensions/e4_test.go +++ b/field/babybear/extensions/e4_test.go @@ -1145,20 +1145,17 @@ func TestE4MarshalSetBytesRoundTrip(t *testing.T) { } var y E4 - if !y.SetBytes(b).Equal(&x) { + _, err := y.SetBytes(b) + require.NoError(t, err) + if !y.Equal(&x) { t.Fatal("SetBytes(Marshal(x)) != x") } } for _, n := range []int{0, BytesE4 - 1, BytesE4 + 1} { - func() { - defer func() { - if recover() == nil { - t.Fatalf("SetBytes did not panic on %d bytes", n) - } - }() - var z E4 - z.SetBytes(make([]byte, n)) - }() + var z E4 + if _, err := z.SetBytes(make([]byte, n)); err == nil { + t.Fatalf("SetBytes did not fail on %d bytes", n) + } } } diff --git a/field/babybear/extensions/e6.go b/field/babybear/extensions/e6.go index c4af4e1847..38e5f41c75 100644 --- a/field/babybear/extensions/e6.go +++ b/field/babybear/extensions/e6.go @@ -6,6 +6,7 @@ package extensions import ( + "fmt" "math/big" "math/bits" "unsafe" @@ -119,10 +120,10 @@ func (z *E6) Marshal() []byte { } // SetBytes sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytes, and returns z. -// It panics if len(b) != BytesE6. -func (z *E6) SetBytes(b []byte) *E6 { +// It returns an error if len(b) != BytesE6. +func (z *E6) SetBytes(b []byte) (*E6, error) { if len(b) != BytesE6 { - panic("E6.SetBytes: invalid input length") + return nil, fmt.Errorf("E6.SetBytes: got %d bytes, expected %d", len(b), BytesE6) } z.B0.A0.SetBytes(b[0*fr.Bytes : 1*fr.Bytes]) z.B0.A1.SetBytes(b[1*fr.Bytes : 2*fr.Bytes]) @@ -130,7 +131,7 @@ func (z *E6) SetBytes(b []byte) *E6 { z.B1.A1.SetBytes(b[3*fr.Bytes : 4*fr.Bytes]) z.B2.A0.SetBytes(b[4*fr.Bytes : 5*fr.Bytes]) z.B2.A1.SetBytes(b[5*fr.Bytes : 6*fr.Bytes]) - return z + return z, nil } // MulByElement multiplies an element in E6 by an element in fr. diff --git a/field/babybear/extensions/e6_test.go b/field/babybear/extensions/e6_test.go index 51bf04fa8c..2bf902dd5c 100644 --- a/field/babybear/extensions/e6_test.go +++ b/field/babybear/extensions/e6_test.go @@ -11,6 +11,8 @@ import ( "testing" fr "github.com/consensys/gnark-crypto/field/babybear" + "github.com/stretchr/testify/require" + "github.com/leanovate/gopter" "github.com/leanovate/gopter/prop" ) @@ -378,20 +380,17 @@ func TestE6MarshalSetBytesRoundTrip(t *testing.T) { } var y E6 - if !y.SetBytes(b).Equal(&x) { + _, err := y.SetBytes(b) + require.NoError(t, err) + if !y.Equal(&x) { t.Fatal("SetBytes(Marshal(x)) != x") } } for _, n := range []int{0, BytesE6 - 1, BytesE6 + 1} { - func() { - defer func() { - if recover() == nil { - t.Fatalf("SetBytes did not panic on %d bytes", n) - } - }() - var z E6 - z.SetBytes(make([]byte, n)) - }() + var z E6 + if _, err := z.SetBytes(make([]byte, n)); err == nil { + t.Fatalf("SetBytes did not fail on %d bytes", n) + } } } diff --git a/field/goldilocks/extensions/e2.go b/field/goldilocks/extensions/e2.go index d673316f42..a4079c0634 100644 --- a/field/goldilocks/extensions/e2.go +++ b/field/goldilocks/extensions/e2.go @@ -6,6 +6,7 @@ package extensions import ( + "fmt" "math/big" fr "github.com/consensys/gnark-crypto/field/goldilocks" @@ -100,14 +101,14 @@ func (z *E2) Marshal() []byte { } // SetBytes sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytes, and returns z. -// It panics if len(b) != BytesE2. -func (z *E2) SetBytes(b []byte) *E2 { +// It returns an error if len(b) != BytesE2. +func (z *E2) SetBytes(b []byte) (*E2, error) { if len(b) != BytesE2 { - panic("E2.SetBytes: invalid input length") + return nil, fmt.Errorf("E2.SetBytes: got %d bytes, expected %d", len(b), BytesE2) } z.A0.SetBytes(b[0*fr.Bytes : 1*fr.Bytes]) z.A1.SetBytes(b[1*fr.Bytes : 2*fr.Bytes]) - return z + return z, nil } // SetRandom sets a0 and a1 to random values diff --git a/field/goldilocks/extensions/e2_test.go b/field/goldilocks/extensions/e2_test.go index 0cac5a0f66..e14bf22aa7 100644 --- a/field/goldilocks/extensions/e2_test.go +++ b/field/goldilocks/extensions/e2_test.go @@ -11,6 +11,8 @@ import ( "math/big" "testing" + "github.com/stretchr/testify/require" + fr "github.com/consensys/gnark-crypto/field/goldilocks" "github.com/leanovate/gopter" "github.com/leanovate/gopter/prop" @@ -556,20 +558,17 @@ func TestE2MarshalSetBytesRoundTrip(t *testing.T) { } var y E2 - if !y.SetBytes(b).Equal(&x) { + _, err := y.SetBytes(b) + require.NoError(t, err) + if !y.Equal(&x) { t.Fatal("SetBytes(Marshal(x)) != x") } } for _, n := range []int{0, BytesE2 - 1, BytesE2 + 1} { - func() { - defer func() { - if recover() == nil { - t.Fatalf("SetBytes did not panic on %d bytes", n) - } - }() - var z E2 - z.SetBytes(make([]byte, n)) - }() + var z E2 + if _, err := z.SetBytes(make([]byte, n)); err == nil { + t.Fatalf("SetBytes did not fail on %d bytes", n) + } } } diff --git a/field/koalabear/extensions/e2.go b/field/koalabear/extensions/e2.go index e2f27d899c..b6d18b2cb8 100644 --- a/field/koalabear/extensions/e2.go +++ b/field/koalabear/extensions/e2.go @@ -6,6 +6,7 @@ package extensions import ( + "fmt" "math/big" fr "github.com/consensys/gnark-crypto/field/koalabear" @@ -100,14 +101,14 @@ func (z *E2) Marshal() []byte { } // SetBytes sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytes, and returns z. -// It panics if len(b) != BytesE2. -func (z *E2) SetBytes(b []byte) *E2 { +// It returns an error if len(b) != BytesE2. +func (z *E2) SetBytes(b []byte) (*E2, error) { if len(b) != BytesE2 { - panic("E2.SetBytes: invalid input length") + return nil, fmt.Errorf("E2.SetBytes: got %d bytes, expected %d", len(b), BytesE2) } z.A0.SetBytes(b[0*fr.Bytes : 1*fr.Bytes]) z.A1.SetBytes(b[1*fr.Bytes : 2*fr.Bytes]) - return z + return z, nil } // SetRandom sets a0 and a1 to random values diff --git a/field/koalabear/extensions/e2_test.go b/field/koalabear/extensions/e2_test.go index 048a52a055..c2429fb21c 100644 --- a/field/koalabear/extensions/e2_test.go +++ b/field/koalabear/extensions/e2_test.go @@ -11,6 +11,8 @@ import ( "math/big" "testing" + "github.com/stretchr/testify/require" + fr "github.com/consensys/gnark-crypto/field/koalabear" "github.com/leanovate/gopter" "github.com/leanovate/gopter/prop" @@ -573,20 +575,17 @@ func TestE2MarshalSetBytesRoundTrip(t *testing.T) { } var y E2 - if !y.SetBytes(b).Equal(&x) { + _, err := y.SetBytes(b) + require.NoError(t, err) + if !y.Equal(&x) { t.Fatal("SetBytes(Marshal(x)) != x") } } for _, n := range []int{0, BytesE2 - 1, BytesE2 + 1} { - func() { - defer func() { - if recover() == nil { - t.Fatalf("SetBytes did not panic on %d bytes", n) - } - }() - var z E2 - z.SetBytes(make([]byte, n)) - }() + var z E2 + if _, err := z.SetBytes(make([]byte, n)); err == nil { + t.Fatalf("SetBytes did not fail on %d bytes", n) + } } } diff --git a/field/koalabear/extensions/e4.go b/field/koalabear/extensions/e4.go index 3622804b3a..9a2a4b2436 100644 --- a/field/koalabear/extensions/e4.go +++ b/field/koalabear/extensions/e4.go @@ -6,6 +6,7 @@ package extensions import ( + "fmt" "math/big" "math/bits" @@ -113,16 +114,16 @@ func (z *E4) Marshal() []byte { } // SetBytes sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytes, and returns z. -// It panics if len(b) != BytesE4. -func (z *E4) SetBytes(b []byte) *E4 { +// It returns an error if len(b) != BytesE4. +func (z *E4) SetBytes(b []byte) (*E4, error) { if len(b) != BytesE4 { - panic("E4.SetBytes: invalid input length") + return nil, fmt.Errorf("E4.SetBytes: got %d bytes, expected %d", len(b), BytesE4) } z.B0.A0.SetBytes(b[0*fr.Bytes : 1*fr.Bytes]) z.B0.A1.SetBytes(b[1*fr.Bytes : 2*fr.Bytes]) z.B1.A0.SetBytes(b[2*fr.Bytes : 3*fr.Bytes]) z.B1.A1.SetBytes(b[3*fr.Bytes : 4*fr.Bytes]) - return z + return z, nil } // Lift sets the B0.A0 component of z to v diff --git a/field/koalabear/extensions/e4_test.go b/field/koalabear/extensions/e4_test.go index 6394b3d676..46f4bf75a1 100644 --- a/field/koalabear/extensions/e4_test.go +++ b/field/koalabear/extensions/e4_test.go @@ -1145,20 +1145,17 @@ func TestE4MarshalSetBytesRoundTrip(t *testing.T) { } var y E4 - if !y.SetBytes(b).Equal(&x) { + _, err := y.SetBytes(b) + require.NoError(t, err) + if !y.Equal(&x) { t.Fatal("SetBytes(Marshal(x)) != x") } } for _, n := range []int{0, BytesE4 - 1, BytesE4 + 1} { - func() { - defer func() { - if recover() == nil { - t.Fatalf("SetBytes did not panic on %d bytes", n) - } - }() - var z E4 - z.SetBytes(make([]byte, n)) - }() + var z E4 + if _, err := z.SetBytes(make([]byte, n)); err == nil { + t.Fatalf("SetBytes did not fail on %d bytes", n) + } } } diff --git a/field/koalabear/extensions/e6.go b/field/koalabear/extensions/e6.go index ad7f664ed9..6a2c755d4a 100644 --- a/field/koalabear/extensions/e6.go +++ b/field/koalabear/extensions/e6.go @@ -6,6 +6,7 @@ package extensions import ( + "fmt" "math/big" "math/bits" "unsafe" @@ -120,10 +121,10 @@ func (z *E6) Marshal() []byte { } // SetBytes sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytes, and returns z. -// It panics if len(b) != BytesE6. -func (z *E6) SetBytes(b []byte) *E6 { +// It returns an error if len(b) != BytesE6. +func (z *E6) SetBytes(b []byte) (*E6, error) { if len(b) != BytesE6 { - panic("E6.SetBytes: invalid input length") + return nil, fmt.Errorf("E6.SetBytes: got %d bytes, expected %d", len(b), BytesE6) } z.B0.A0.SetBytes(b[0*fr.Bytes : 1*fr.Bytes]) z.B0.A1.SetBytes(b[1*fr.Bytes : 2*fr.Bytes]) @@ -131,7 +132,7 @@ func (z *E6) SetBytes(b []byte) *E6 { z.B1.A1.SetBytes(b[3*fr.Bytes : 4*fr.Bytes]) z.B2.A0.SetBytes(b[4*fr.Bytes : 5*fr.Bytes]) z.B2.A1.SetBytes(b[5*fr.Bytes : 6*fr.Bytes]) - return z + return z, nil } // MulByElement multiplies an element in E6 by an element in fr. diff --git a/field/koalabear/extensions/e6_test.go b/field/koalabear/extensions/e6_test.go index 4fd52d1017..27129696df 100644 --- a/field/koalabear/extensions/e6_test.go +++ b/field/koalabear/extensions/e6_test.go @@ -11,6 +11,8 @@ import ( "testing" fr "github.com/consensys/gnark-crypto/field/koalabear" + "github.com/stretchr/testify/require" + "github.com/leanovate/gopter" "github.com/leanovate/gopter/prop" ) @@ -378,20 +380,17 @@ func TestE6MarshalSetBytesRoundTrip(t *testing.T) { } var y E6 - if !y.SetBytes(b).Equal(&x) { + _, err := y.SetBytes(b) + require.NoError(t, err) + if !y.Equal(&x) { t.Fatal("SetBytes(Marshal(x)) != x") } } for _, n := range []int{0, BytesE6 - 1, BytesE6 + 1} { - func() { - defer func() { - if recover() == nil { - t.Fatalf("SetBytes did not panic on %d bytes", n) - } - }() - var z E6 - z.SetBytes(make([]byte, n)) - }() + var z E6 + if _, err := z.SetBytes(make([]byte, n)); err == nil { + t.Fatalf("SetBytes did not fail on %d bytes", n) + } } } diff --git a/field/koalabear/extensions/polynomial/multilin.go b/field/koalabear/extensions/polynomial/multilin_e6.go similarity index 85% rename from field/koalabear/extensions/polynomial/multilin.go rename to field/koalabear/extensions/polynomial/multilin_e6.go index d0c45379fa..929478d873 100644 --- a/field/koalabear/extensions/polynomial/multilin.go +++ b/field/koalabear/extensions/polynomial/multilin_e6.go @@ -12,14 +12,14 @@ import ( "github.com/consensys/gnark-crypto/utils" ) -// MultiLin tracks the values of a (dense i.e. not sparse) multilinear polynomial +// MultiLinE6 tracks the values of a (dense i.e. not sparse) multilinear polynomial // The variables are X₁ through Xₙ where n = log(len(.)) // .[∑ᵢ 2ⁱ⁻¹ bₙ₋ᵢ] = the polynomial evaluated at (b₁, b₂, ..., bₙ) // It is understood that any hypercube evaluation can be extrapolated to a multilinear polynomial -type MultiLin []extensions.E6 +type MultiLinE6 []extensions.E6 // Fold is partial evaluation function k[X₁, X₂, ..., Xₙ] → k[X₂, ..., Xₙ] by setting X₁=r -func (m *MultiLin) Fold(r extensions.E6) { +func (m *MultiLinE6) Fold(r extensions.E6) { mid := len(*m) / 2 bottom, top := (*m)[:mid], (*m)[mid:] @@ -40,7 +40,7 @@ func (m *MultiLin) Fold(r extensions.E6) { *m = (*m)[:mid] } -func (m *MultiLin) FoldParallel(r extensions.E6) utils.Task { +func (m *MultiLinE6) FoldParallel(r extensions.E6) utils.Task { mid := len(*m) / 2 bottom, top := (*m)[:mid], (*m)[mid:] @@ -57,7 +57,7 @@ func (m *MultiLin) FoldParallel(r extensions.E6) utils.Task { } } -func (m MultiLin) Sum() extensions.E6 { +func (m MultiLinE6) Sum() extensions.E6 { s := m[0] for i := 1; i < len(m); i++ { s.Add(&s, &m[i]) @@ -65,7 +65,7 @@ func (m MultiLin) Sum() extensions.E6 { return s } -func _clone(m MultiLin, p *Pool) MultiLin { +func _cloneE6(m MultiLinE6, p *PoolE6) MultiLinE6 { if p == nil { return m.Clone() } else { @@ -73,7 +73,7 @@ func _clone(m MultiLin, p *Pool) MultiLin { } } -func _dump(m MultiLin, p *Pool) { +func _dumpE6(m MultiLinE6, p *PoolE6) { if p != nil { p.Dump(m) } @@ -81,9 +81,9 @@ func _dump(m MultiLin, p *Pool) { // Evaluate extrapolate the value of the multilinear polynomial corresponding to m // on the given coordinates -func (m MultiLin) Evaluate(coordinates []extensions.E6, p *Pool) extensions.E6 { +func (m MultiLinE6) Evaluate(coordinates []extensions.E6, p *PoolE6) extensions.E6 { // Folding is a mutating operation - bkCopy := _clone(m, p) + bkCopy := _cloneE6(m, p) // Evaluate step by step through repeated folding (i.e. evaluation at the first remaining variable) for _, r := range coordinates { @@ -92,7 +92,7 @@ func (m MultiLin) Evaluate(coordinates []extensions.E6, p *Pool) extensions.E6 { result := bkCopy[0] - _dump(bkCopy, p) + _dumpE6(bkCopy, p) return result } @@ -100,14 +100,14 @@ func (m MultiLin) Evaluate(coordinates []extensions.E6, p *Pool) extensions.E6 { // Both multilinear interpolation and sumcheck require folding an underlying // array, but folding changes the array. To do both one requires a deep copy // of the bookkeeping table. -func (m MultiLin) Clone() MultiLin { - res := make(MultiLin, len(m)) +func (m MultiLinE6) Clone() MultiLinE6 { + res := make(MultiLinE6, len(m)) copy(res, m) return res } // Add two bookKeepingTables -func (m *MultiLin) Add(left, right MultiLin) { +func (m *MultiLinE6) Add(left, right MultiLinE6) { size := len(left) // Check that left and right have the same size if len(right) != size || len(*m) != size { @@ -120,7 +120,7 @@ func (m *MultiLin) Add(left, right MultiLin) { } } -// EvalEq computes Eq(q₁, ... , qₙ, h₁, ... , hₙ) = Π₁ⁿ Eq(qᵢ, hᵢ) +// EvalEqE6 computes Eq(q₁, ... , qₙ, h₁, ... , hₙ) = Π₁ⁿ Eq(qᵢ, hᵢ) // where Eq(x,y) = xy + (1-x)(1-y) = 1 - x - y + xy + xy interpolates // // _________________ @@ -135,7 +135,7 @@ func (m *MultiLin) Add(left, right MultiLin) { // // In other words the polynomial evaluated here is the multilinear extrapolation of // one that evaluates to q' == h' for vectors q', h' of binary values -func EvalEq(q, h []extensions.E6) extensions.E6 { +func EvalEqE6(q, h []extensions.E6) extensions.E6 { var res, nxt, one, sum extensions.E6 one.SetOne() for i := range len(q) { @@ -155,7 +155,7 @@ func EvalEq(q, h []extensions.E6) extensions.E6 { } // Eq sets m to the representation of the polynomial Eq(q₁, ..., qₙ, *, ..., *) × m[0] -func (m *MultiLin) Eq(q []extensions.E6) { +func (m *MultiLinE6) Eq(q []extensions.E6) { n := len(q) if len(*m) != 1<= 0; i-- { @@ -33,14 +33,14 @@ func (p *Polynomial) Eval(v *extensions.E6) extensions.E6 { } // Clone returns a copy of the polynomial -func (p *Polynomial) Clone() Polynomial { - _p := make(Polynomial, len(*p)) +func (p *PolynomialE6) Clone() PolynomialE6 { + _p := make(PolynomialE6, len(*p)) copy(_p, *p) return _p } // Set to another polynomial -func (p *Polynomial) Set(p1 Polynomial) { +func (p *PolynomialE6) Set(p1 PolynomialE6) { if len(*p) != len(p1) { *p = p1.Clone() return @@ -52,30 +52,30 @@ func (p *Polynomial) Set(p1 Polynomial) { } // AddConstantInPlace adds a constant to the polynomial, modifying p -func (p *Polynomial) AddConstantInPlace(c *extensions.E6) { +func (p *PolynomialE6) AddConstantInPlace(c *extensions.E6) { for i := range len(*p) { (*p)[i].Add(&(*p)[i], c) } } // SubConstantInPlace subs a constant to the polynomial, modifying p -func (p *Polynomial) SubConstantInPlace(c *extensions.E6) { +func (p *PolynomialE6) SubConstantInPlace(c *extensions.E6) { for i := range len(*p) { (*p)[i].Sub(&(*p)[i], c) } } // ScaleInPlace multiplies p by v, modifying p -func (p *Polynomial) ScaleInPlace(c *extensions.E6) { +func (p *PolynomialE6) ScaleInPlace(c *extensions.E6) { for i := range len(*p) { (*p)[i].Mul(&(*p)[i], c) } } // Scale multiplies p0 by v, storing the result in p -func (p *Polynomial) Scale(c *extensions.E6, p0 Polynomial) { +func (p *PolynomialE6) Scale(c *extensions.E6, p0 PolynomialE6) { if len(*p) != len(p0) { - *p = make(Polynomial, len(p0)) + *p = make(PolynomialE6, len(p0)) } for i := range len(p0) { (*p)[i].Mul(c, &p0[i]) @@ -84,7 +84,7 @@ func (p *Polynomial) Scale(c *extensions.E6, p0 Polynomial) { // Add adds p1 to p2 // This function allocates a new slice unless p == p1 or p == p2 -func (p *Polynomial) Add(p1, p2 Polynomial) *Polynomial { +func (p *PolynomialE6) Add(p1, p2 PolynomialE6) *PolynomialE6 { bigger := p1 smaller := p2 @@ -107,7 +107,7 @@ func (p *Polynomial) Add(p1, p2 Polynomial) *Polynomial { return p } - res := make(Polynomial, len(bigger)) + res := make(PolynomialE6, len(bigger)) copy(res, bigger) for i := range len(smaller) { res[i].Add(&res[i], &smaller[i]) @@ -118,7 +118,7 @@ func (p *Polynomial) Add(p1, p2 Polynomial) *Polynomial { // Sub subtracts p2 from p1 // TODO make interface more consistent with Add -func (p *Polynomial) Sub(p1, p2 Polynomial) *Polynomial { +func (p *PolynomialE6) Sub(p1, p2 PolynomialE6) *PolynomialE6 { if len(p1) != len(p2) || len(p2) != len(*p) { return nil } @@ -129,7 +129,7 @@ func (p *Polynomial) Sub(p1, p2 Polynomial) *Polynomial { } // Equal checks equality between two polynomials -func (p *Polynomial) Equal(p1 Polynomial) bool { +func (p *PolynomialE6) Equal(p1 PolynomialE6) bool { if (*p == nil) != (p1 == nil) { return false } @@ -147,26 +147,26 @@ func (p *Polynomial) Equal(p1 Polynomial) bool { return true } -func (p Polynomial) SetZero() { +func (p PolynomialE6) SetZero() { for i := range len(p) { p[i].SetZero() } } -// InterpolateOnRange maps vector v to polynomial f +// InterpolateOnRangeE6 maps vector v to polynomial f // such that f(i) = v[i] for 0 ≤ i < len(v). // len(f) = len(v) and deg(f) ≤ len(v) - 1 -func InterpolateOnRange(v []extensions.E6) Polynomial { +func InterpolateOnRangeE6(v []extensions.E6) PolynomialE6 { nEvals := uint8(len(v)) if int(nEvals) != len(v) { panic("interpolation method too inefficient for nEvals > 255") } - lagrange := getLagrangeBasis(nEvals) + lagrange := getLagrangeBasisE6(nEvals) - var res Polynomial + var res PolynomialE6 res.Scale(&v[0], lagrange[0]) - temp := make(Polynomial, nEvals) + temp := make(PolynomialE6, nEvals) for i := uint8(1); i < nEvals; i++ { temp.Scale(&v[i], lagrange[i]) @@ -176,39 +176,39 @@ func InterpolateOnRange(v []extensions.E6) Polynomial { return res } -// lagrange bases used by InterpolateOnRange -var lagrangeBasis sync.Map +// lagrange bases used by InterpolateOnRangeE6 +var lagrangeBasisE6 sync.Map -func getLagrangeBasis(domainSize uint8) []Polynomial { - if res, ok := lagrangeBasis.Load(domainSize); ok { - return res.([]Polynomial) +func getLagrangeBasisE6(domainSize uint8) []PolynomialE6 { + if res, ok := lagrangeBasisE6.Load(domainSize); ok { + return res.([]PolynomialE6) } // not found. compute - var res []Polynomial + var res []PolynomialE6 if domainSize >= 2 { - res = computeLagrangeBasis(domainSize) + res = computeLagrangeBasisE6(domainSize) } else if domainSize == 1 { - res = []Polynomial{make(Polynomial, 1)} + res = []PolynomialE6{make(PolynomialE6, 1)} res[0][0].SetOne() } - lagrangeBasis.Store(domainSize, res) + lagrangeBasisE6.Store(domainSize, res) return res } -// computeLagrangeBasis precomputes in explicit coefficient form for each 0 ≤ l < domainSize the polynomial +// computeLagrangeBasisE6 precomputes in explicit coefficient form for each 0 ≤ l < domainSize the polynomial // pₗ := X (X-1) ... (X-l-1) (X-l+1) ... (X - domainSize + 1) / ( l (l-1) ... 2 (-1) ... (l - domainSize +1) ) // Note that pₗ(l) = 1 and pₗ(n) = 0 if 0 ≤ l < domainSize, n ≠ l -func computeLagrangeBasis(domainSize uint8) []Polynomial { +func computeLagrangeBasisE6(domainSize uint8) []PolynomialE6 { constTerms := make([]extensions.E6, domainSize) for i := range domainSize { constTerms[i].SetInt64(-int64(i)) } - res := make([]Polynomial, domainSize) - multScratch := make(Polynomial, domainSize-1) + res := make([]PolynomialE6, domainSize) + multScratch := make(PolynomialE6, domainSize-1) // compute pₗ for l := range domainSize { @@ -220,7 +220,7 @@ func computeLagrangeBasis(domainSize uint8) []Polynomial { continue } if d == 0 { - res[l] = make(Polynomial, domainSize) + res[l] = make(PolynomialE6, domainSize) res[l][domainSize-2] = constTerms[i] res[l][domainSize-1].SetOne() } else { diff --git a/field/koalabear/extensions/polynomial/polynomial_test.go b/field/koalabear/extensions/polynomial/polynomial_e6_test.go similarity index 84% rename from field/koalabear/extensions/polynomial/polynomial_test.go rename to field/koalabear/extensions/polynomial/polynomial_e6_test.go index 711574f4ec..27862f1121 100644 --- a/field/koalabear/extensions/polynomial/polynomial_test.go +++ b/field/koalabear/extensions/polynomial/polynomial_e6_test.go @@ -16,10 +16,10 @@ import ( "github.com/stretchr/testify/assert" ) -func TestPolynomialEval(t *testing.T) { +func TestPolynomialEvalE6(t *testing.T) { // build polynomial - f := make(Polynomial, 20) + f := make(PolynomialE6, 20) for i := range 20 { f[i].SetOne() } @@ -47,10 +47,10 @@ func TestPolynomialEval(t *testing.T) { } } -func TestPolynomialAddConstantInPlace(t *testing.T) { +func TestPolynomialAddConstantInPlaceE6(t *testing.T) { // build polynomial - f := make(Polynomial, 20) + f := make(PolynomialE6, 20) for i := range 20 { f[i].SetOne() } @@ -73,10 +73,10 @@ func TestPolynomialAddConstantInPlace(t *testing.T) { } } -func TestPolynomialSubConstantInPlace(t *testing.T) { +func TestPolynomialSubConstantInPlaceE6(t *testing.T) { // build polynomial - f := make(Polynomial, 20) + f := make(PolynomialE6, 20) for i := range 20 { f[i].SetOne() } @@ -99,10 +99,10 @@ func TestPolynomialSubConstantInPlace(t *testing.T) { } } -func TestPolynomialScaleInPlace(t *testing.T) { +func TestPolynomialScaleInPlaceE6(t *testing.T) { // build polynomial - f := make(Polynomial, 20) + f := make(PolynomialE6, 20) for i := range 20 { f[i].SetOne() } @@ -123,17 +123,17 @@ func TestPolynomialScaleInPlace(t *testing.T) { } -func TestPolynomialAdd(t *testing.T) { +func TestPolynomialAddE6(t *testing.T) { // build unbalanced polynomials - f1 := make(Polynomial, 20) - f1Backup := make(Polynomial, 20) + f1 := make(PolynomialE6, 20) + f1Backup := make(PolynomialE6, 20) for i := range 20 { f1[i].SetOne() f1Backup[i].SetOne() } - f2 := make(Polynomial, 10) - f2Backup := make(Polynomial, 10) + f2 := make(PolynomialE6, 10) + f2Backup := make(PolynomialE6, 10) for i := range 10 { f2[i].SetOne() f2Backup[i].SetOne() @@ -143,7 +143,7 @@ func TestPolynomialAdd(t *testing.T) { var one, two extensions.E6 one.SetOne() two.Double(&one) - expectedSum := make(Polynomial, 20) + expectedSum := make(PolynomialE6, 20) for i := range 10 { expectedSum[i].Set(&two) } @@ -152,7 +152,7 @@ func TestPolynomialAdd(t *testing.T) { } // caller is empty - var g Polynomial + var g PolynomialE6 g.Add(f1, f2) if !g.Equal(expectedSum) { t.Fatal("add polynomials fails") @@ -200,10 +200,10 @@ func TestPolynomialAdd(t *testing.T) { } } -func TestPrecomputeLagrange(t *testing.T) { +func TestPrecomputeLagrangeE6(t *testing.T) { testForDomainSize := func(domainSize uint8) bool { - polys := computeLagrangeBasis(domainSize) + polys := computeLagrangeBasisE6(domainSize) for l := range domainSize { for i := range domainSize { @@ -237,9 +237,9 @@ func TestPrecomputeLagrange(t *testing.T) { properties.TestingRun(t, gopter.ConsoleReporter(false)) } -func TestLagrangeCache(t *testing.T) { +func TestLagrangeCacheE6(t *testing.T) { for _, i := range []int{5, 2, 8, 4, 6, 3, 0} { - b := getLagrangeBasis(uint8(i)) - assert.Equal(t, b, getLagrangeBasis(uint8(i))) // second call must yield the same result + b := getLagrangeBasisE6(uint8(i)) + assert.Equal(t, b, getLagrangeBasisE6(uint8(i))) // second call must yield the same result } } diff --git a/field/koalabear/extensions/polynomial/pool.go b/field/koalabear/extensions/polynomial/pool_e6.go similarity index 67% rename from field/koalabear/extensions/polynomial/pool.go rename to field/koalabear/extensions/polynomial/pool_e6.go index 3f246202f7..529e80630e 100644 --- a/field/koalabear/extensions/polynomial/pool.go +++ b/field/koalabear/extensions/polynomial/pool_e6.go @@ -20,38 +20,38 @@ import ( // WARNING: This is not thread safe TODO: Make sure that is not a problem // TODO: There is a lot of "unsafe" memory management here and needs to be vetted thoroughly -type sizedPool struct { +type sizedPoolE6 struct { maxN int pool sync.Pool - stats poolStats + stats poolStatsE6 } -type inUseData struct { +type inUseDataE6 struct { allocatedFor []uintptr - pool *sizedPool + pool *sizedPoolE6 } -type Pool struct { +type PoolE6 struct { //lock sync.Mutex inUse sync.Map - subPools []sizedPool + subPools []sizedPoolE6 } -func (p *sizedPool) get(n int) *extensions.E6 { +func (p *sizedPoolE6) get(n int) *extensions.E6 { p.stats.make(n) return p.pool.Get().(*extensions.E6) } -func (p *sizedPool) put(ptr *extensions.E6) { +func (p *sizedPoolE6) put(ptr *extensions.E6) { p.stats.dump() p.pool.Put(ptr) } -func NewPool(maxN ...int) (pool Pool) { +func NewPoolE6(maxN ...int) (pool PoolE6) { sort.Ints(maxN) - pool = Pool{ - subPools: make([]sizedPool, len(maxN)), + pool = PoolE6{ + subPools: make([]sizedPoolE6, len(maxN)), } for i := range pool.subPools { @@ -60,14 +60,14 @@ func NewPool(maxN ...int) (pool Pool) { subPool.pool = sync.Pool{ New: func() any { subPool.stats.Allocated++ - return getDataPointer(make([]extensions.E6, 0, subPool.maxN)) + return getDataPointerE6(make([]extensions.E6, 0, subPool.maxN)) }, } } return } -func (p *Pool) findCorrespondingPool(n int) *sizedPool { +func (p *PoolE6) findCorrespondingPool(n int) *sizedPoolE6 { poolI := 0 for poolI < len(p.subPools) && n > p.subPools[poolI].maxN { poolI++ @@ -75,7 +75,7 @@ func (p *Pool) findCorrespondingPool(n int) *sizedPool { return &p.subPools[poolI] // out of bounds error here would mean that n is too large } -func (p *Pool) Make(n int) []extensions.E6 { +func (p *PoolE6) Make(n int) []extensions.E6 { pool := p.findCorrespondingPool(n) ptr := pool.get(n) p.addInUse(ptr, pool) @@ -83,52 +83,52 @@ func (p *Pool) Make(n int) []extensions.E6 { } // Dump dumps a set of polynomials into the pool -func (p *Pool) Dump(slices ...[]extensions.E6) { +func (p *PoolE6) Dump(slices ...[]extensions.E6) { for _, slice := range slices { - ptr := getDataPointer(slice) + ptr := getDataPointerE6(slice) if metadata, ok := p.inUse.Load(ptr); ok { p.inUse.Delete(ptr) - metadata.(inUseData).pool.put(ptr) + metadata.(inUseDataE6).pool.put(ptr) } else { panic("attempting to dump a slice not created by the pool") } } } -func (p *Pool) addInUse(ptr *extensions.E6, pool *sizedPool) { +func (p *PoolE6) addInUse(ptr *extensions.E6, pool *sizedPoolE6) { pcs := make([]uintptr, 2) n := runtime.Callers(3, pcs) if prevPcs, ok := p.inUse.Load(ptr); ok { // TODO: remove if unnecessary for security - panic(fmt.Errorf("re-allocated non-dumped slice, previously allocated at %v", runtime.CallersFrames(prevPcs.(inUseData).allocatedFor))) + panic(fmt.Errorf("re-allocated non-dumped slice, previously allocated at %v", runtime.CallersFrames(prevPcs.(inUseDataE6).allocatedFor))) } - p.inUse.Store(ptr, inUseData{ + p.inUse.Store(ptr, inUseDataE6{ allocatedFor: pcs[:n], pool: pool, }) } -func printFrame(frame runtime.Frame) { +func printFrameE6(frame runtime.Frame) { fmt.Printf("\t%s line %d, function %s\n", frame.File, frame.Line, frame.Function) } -func (p *Pool) printInUse() { +func (p *PoolE6) printInUse() { fmt.Println("slices never dumped allocated at:") p.inUse.Range(func(_, pcs any) bool { fmt.Println("-------------------------") var frame runtime.Frame - frames := runtime.CallersFrames(pcs.(inUseData).allocatedFor) + frames := runtime.CallersFrames(pcs.(inUseDataE6).allocatedFor) more := true for more { frame, more = frames.Next() - printFrame(frame) + printFrameE6(frame) } return true }) } -type poolStats struct { +type poolStatsE6 struct { Used int Allocated int ReuseRate float64 @@ -137,12 +137,12 @@ type poolStats struct { SmallestNUsed int } -type poolsStats struct { - SubPools []poolStats +type poolsStatsE6 struct { + SubPools []poolStatsE6 InUse int } -func (s *poolStats) make(n int) { +func (s *poolStatsE6) make(n int) { s.Used++ s.InUse++ if n > s.GreatestNUsed { @@ -153,21 +153,21 @@ func (s *poolStats) make(n int) { } } -func (s *poolStats) dump() { +func (s *poolStatsE6) dump() { s.InUse-- } -func (s *poolStats) finalize() { +func (s *poolStatsE6) finalize() { s.ReuseRate = float64(s.Used) / float64(s.Allocated) } -func getDataPointer(slice []extensions.E6) *extensions.E6 { +func getDataPointerE6(slice []extensions.E6) *extensions.E6 { return (*extensions.E6)(unsafe.SliceData(slice)) } -func (p *Pool) PrintPoolStats() { +func (p *PoolE6) PrintPoolStats() { InUse := 0 - subStats := make([]poolStats, len(p.subPools)) + subStats := make([]poolStatsE6, len(p.subPools)) for i := range p.subPools { subPool := &p.subPools[i] subPool.stats.finalize() @@ -175,7 +175,7 @@ func (p *Pool) PrintPoolStats() { InUse += subPool.stats.InUse } - stats := poolsStats{ + stats := poolsStatsE6{ SubPools: subStats, InUse: InUse, } @@ -184,7 +184,7 @@ func (p *Pool) PrintPoolStats() { p.printInUse() } -func (p *Pool) Clone(slice []extensions.E6) []extensions.E6 { +func (p *PoolE6) Clone(slice []extensions.E6) []extensions.E6 { res := p.Make(len(slice)) copy(res, slice) return res diff --git a/field/mamabear/extensions/e3.go b/field/mamabear/extensions/e3.go index 84fd32b301..490342cf88 100644 --- a/field/mamabear/extensions/e3.go +++ b/field/mamabear/extensions/e3.go @@ -6,6 +6,7 @@ package extensions import ( + "fmt" "math/big" "sync" @@ -69,15 +70,15 @@ func (z *E3) Marshal() []byte { } // SetBytes sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytes, and returns z. -// It panics if len(b) != BytesE3. -func (z *E3) SetBytes(b []byte) *E3 { +// It returns an error if len(b) != BytesE3. +func (z *E3) SetBytes(b []byte) (*E3, error) { if len(b) != BytesE3 { - panic("E3.SetBytes: invalid input length") + return nil, fmt.Errorf("E3.SetBytes: got %d bytes, expected %d", len(b), BytesE3) } z.A0.SetBytes(b[0*fr.Bytes : 1*fr.Bytes]) z.A1.SetBytes(b[1*fr.Bytes : 2*fr.Bytes]) z.A2.SetBytes(b[2*fr.Bytes : 3*fr.Bytes]) - return z + return z, nil } // IsZero reports whether z is zero. diff --git a/field/mamabear/extensions/e3_test.go b/field/mamabear/extensions/e3_test.go index bd8fe05104..7c381f14d0 100644 --- a/field/mamabear/extensions/e3_test.go +++ b/field/mamabear/extensions/e3_test.go @@ -10,6 +10,8 @@ import ( "math/big" "testing" + "github.com/stretchr/testify/require" + fr "github.com/consensys/gnark-crypto/field/mamabear" ) @@ -462,20 +464,17 @@ func TestE3MarshalSetBytesRoundTrip(t *testing.T) { } var y E3 - if !y.SetBytes(b).Equal(&x) { + _, err := y.SetBytes(b) + require.NoError(t, err) + if !y.Equal(&x) { t.Fatal("SetBytes(Marshal(x)) != x") } } for _, n := range []int{0, BytesE3 - 1, BytesE3 + 1} { - func() { - defer func() { - if recover() == nil { - t.Fatalf("SetBytes did not panic on %d bytes", n) - } - }() - var z E3 - z.SetBytes(make([]byte, n)) - }() + var z E3 + if _, err := z.SetBytes(make([]byte, n)); err == nil { + t.Fatalf("SetBytes did not fail on %d bytes", n) + } } } diff --git a/internal/generator/config/fields.go b/internal/generator/config/fields.go index 8bc800d8be..23ba42032b 100644 --- a/internal/generator/config/fields.go +++ b/internal/generator/config/fields.go @@ -9,6 +9,10 @@ type Field struct { // HandwrittenVectorASMAMD64 declares that this field ships hand-written // amd64 vector and FFT kernels next to its generated output. HandwrittenVectorASMAMD64 bool + // PolynomialExtensions lists the degrees of the extensions of this field + // over which the polynomial package (Polynomial, MultiLin, Pool, ...) is + // generated. Empty means no polynomial package. + PolynomialExtensions []int } var Fields []Field @@ -23,8 +27,9 @@ func init() { Modulus: "0xFFFFFFFF00000001", }) addField(Field{ - Name: "koalabear", - Modulus: "0x7f000001", + Name: "koalabear", + Modulus: "0x7f000001", + PolynomialExtensions: []int{6}, }) addField(Field{ Name: "babybear", diff --git a/internal/generator/field/config/field_config.go b/internal/generator/field/config/field_config.go index b1176865ec..4021921d4d 100644 --- a/internal/generator/field/config/field_config.go +++ b/internal/generator/field/config/field_config.go @@ -940,3 +940,13 @@ type FieldDependency struct { FieldPackageName string ExtensionDegree int } + +// ExtensionName returns the name of the extension (e.g. "E6"), or the empty +// string for a base field. It is appended to the names of generated free +// functions and types so that several extensions can share a package. +func (d FieldDependency) ExtensionName() string { + if d.ExtensionDegree == 0 { + return "" + } + return fmt.Sprintf("E%d", d.ExtensionDegree) +} diff --git a/internal/generator/field/template/extensions/e2.go.tmpl b/internal/generator/field/template/extensions/e2.go.tmpl index 8d5daf2bf6..114b4a7e34 100644 --- a/internal/generator/field/template/extensions/e2.go.tmpl +++ b/internal/generator/field/template/extensions/e2.go.tmpl @@ -93,14 +93,14 @@ func (z *E2) Marshal() []byte { } // SetBytes sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytes, and returns z. -// It panics if len(b) != BytesE2. -func (z *E2) SetBytes(b []byte) *E2 { +// It returns an error if len(b) != BytesE2. +func (z *E2) SetBytes(b []byte) (*E2, error) { if len(b) != BytesE2 { - panic("E2.SetBytes: invalid input length") + return nil, fmt.Errorf("E2.SetBytes: got %d bytes, expected %d", len(b), BytesE2) } z.A0.SetBytes(b[0*fr.Bytes : 1*fr.Bytes]) z.A1.SetBytes(b[1*fr.Bytes : 2*fr.Bytes]) - return z + return z, nil } // SetRandom sets a0 and a1 to random values diff --git a/internal/generator/field/template/extensions/e2_test.go.tmpl b/internal/generator/field/template/extensions/e2_test.go.tmpl index e01d8e2b8a..b802a7778e 100644 --- a/internal/generator/field/template/extensions/e2_test.go.tmpl +++ b/internal/generator/field/template/extensions/e2_test.go.tmpl @@ -1,4 +1,5 @@ import ( + "github.com/stretchr/testify/require" "crypto/rand" "testing" "math/big" @@ -578,20 +579,17 @@ func TestE2MarshalSetBytesRoundTrip(t *testing.T) { } var y E2 - if !y.SetBytes(b).Equal(&x) { + _, err := y.SetBytes(b) + require.NoError(t, err) + if !y.Equal(&x) { t.Fatal("SetBytes(Marshal(x)) != x") } } for _, n := range []int{0, BytesE2 - 1, BytesE2 + 1} { - func() { - defer func() { - if recover() == nil { - t.Fatalf("SetBytes did not panic on %d bytes", n) - } - }() - var z E2 - z.SetBytes(make([]byte, n)) - }() + var z E2 + if _, err := z.SetBytes(make([]byte, n)); err == nil { + t.Fatalf("SetBytes did not fail on %d bytes", n) + } } } diff --git a/internal/generator/field/template/extensions/e3.go.tmpl b/internal/generator/field/template/extensions/e3.go.tmpl index c6fcb64913..247663c1c7 100644 --- a/internal/generator/field/template/extensions/e3.go.tmpl +++ b/internal/generator/field/template/extensions/e3.go.tmpl @@ -62,15 +62,15 @@ func (z *E3) Marshal() []byte { } // SetBytes sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytes, and returns z. -// It panics if len(b) != BytesE3. -func (z *E3) SetBytes(b []byte) *E3 { +// It returns an error if len(b) != BytesE3. +func (z *E3) SetBytes(b []byte) (*E3, error) { if len(b) != BytesE3 { - panic("E3.SetBytes: invalid input length") + return nil, fmt.Errorf("E3.SetBytes: got %d bytes, expected %d", len(b), BytesE3) } z.A0.SetBytes(b[0*fr.Bytes : 1*fr.Bytes]) z.A1.SetBytes(b[1*fr.Bytes : 2*fr.Bytes]) z.A2.SetBytes(b[2*fr.Bytes : 3*fr.Bytes]) - return z + return z, nil } // IsZero reports whether z is zero. diff --git a/internal/generator/field/template/extensions/e3_test.go.tmpl b/internal/generator/field/template/extensions/e3_test.go.tmpl index 18a93c0625..014eeb51b0 100644 --- a/internal/generator/field/template/extensions/e3_test.go.tmpl +++ b/internal/generator/field/template/extensions/e3_test.go.tmpl @@ -1,4 +1,5 @@ import ( + "github.com/stretchr/testify/require" "math/big" "testing" @@ -454,20 +455,17 @@ func TestE3MarshalSetBytesRoundTrip(t *testing.T) { } var y E3 - if !y.SetBytes(b).Equal(&x) { + _, err := y.SetBytes(b) + require.NoError(t, err) + if !y.Equal(&x) { t.Fatal("SetBytes(Marshal(x)) != x") } } for _, n := range []int{0, BytesE3 - 1, BytesE3 + 1} { - func() { - defer func() { - if recover() == nil { - t.Fatalf("SetBytes did not panic on %d bytes", n) - } - }() - var z E3 - z.SetBytes(make([]byte, n)) - }() + var z E3 + if _, err := z.SetBytes(make([]byte, n)); err == nil { + t.Fatalf("SetBytes did not fail on %d bytes", n) + } } } diff --git a/internal/generator/field/template/extensions/e4.go.tmpl b/internal/generator/field/template/extensions/e4.go.tmpl index 3c6ba0b4b1..6af571dfcd 100644 --- a/internal/generator/field/template/extensions/e4.go.tmpl +++ b/internal/generator/field/template/extensions/e4.go.tmpl @@ -106,16 +106,16 @@ func (z *E4) Marshal() []byte { } // SetBytes sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytes, and returns z. -// It panics if len(b) != BytesE4. -func (z *E4) SetBytes(b []byte) *E4 { +// It returns an error if len(b) != BytesE4. +func (z *E4) SetBytes(b []byte) (*E4, error) { if len(b) != BytesE4 { - panic("E4.SetBytes: invalid input length") + return nil, fmt.Errorf("E4.SetBytes: got %d bytes, expected %d", len(b), BytesE4) } z.B0.A0.SetBytes(b[0*fr.Bytes : 1*fr.Bytes]) z.B0.A1.SetBytes(b[1*fr.Bytes : 2*fr.Bytes]) z.B1.A0.SetBytes(b[2*fr.Bytes : 3*fr.Bytes]) z.B1.A1.SetBytes(b[3*fr.Bytes : 4*fr.Bytes]) - return z + return z, nil } // Lift sets the B0.A0 component of z to v diff --git a/internal/generator/field/template/extensions/e4_test.go.tmpl b/internal/generator/field/template/extensions/e4_test.go.tmpl index ac2879da37..2b6875a1d0 100644 --- a/internal/generator/field/template/extensions/e4_test.go.tmpl +++ b/internal/generator/field/template/extensions/e4_test.go.tmpl @@ -1143,20 +1143,17 @@ func TestE4MarshalSetBytesRoundTrip(t *testing.T) { } var y E4 - if !y.SetBytes(b).Equal(&x) { + _, err := y.SetBytes(b) + require.NoError(t, err) + if !y.Equal(&x) { t.Fatal("SetBytes(Marshal(x)) != x") } } for _, n := range []int{0, BytesE4 - 1, BytesE4 + 1} { - func() { - defer func() { - if recover() == nil { - t.Fatalf("SetBytes did not panic on %d bytes", n) - } - }() - var z E4 - z.SetBytes(make([]byte, n)) - }() + var z E4 + if _, err := z.SetBytes(make([]byte, n)); err == nil { + t.Fatalf("SetBytes did not fail on %d bytes", n) + } } } diff --git a/internal/generator/field/template/extensions/e6.go.tmpl b/internal/generator/field/template/extensions/e6.go.tmpl index f3008945f2..c028861084 100644 --- a/internal/generator/field/template/extensions/e6.go.tmpl +++ b/internal/generator/field/template/extensions/e6.go.tmpl @@ -115,10 +115,10 @@ func (z *E6) Marshal() []byte { } // SetBytes sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytes, and returns z. -// It panics if len(b) != BytesE6. -func (z *E6) SetBytes(b []byte) *E6 { +// It returns an error if len(b) != BytesE6. +func (z *E6) SetBytes(b []byte) (*E6, error) { if len(b) != BytesE6 { - panic("E6.SetBytes: invalid input length") + return nil, fmt.Errorf("E6.SetBytes: got %d bytes, expected %d", len(b), BytesE6) } z.B0.A0.SetBytes(b[0*fr.Bytes : 1*fr.Bytes]) z.B0.A1.SetBytes(b[1*fr.Bytes : 2*fr.Bytes]) @@ -126,7 +126,7 @@ func (z *E6) SetBytes(b []byte) *E6 { z.B1.A1.SetBytes(b[3*fr.Bytes : 4*fr.Bytes]) z.B2.A0.SetBytes(b[4*fr.Bytes : 5*fr.Bytes]) z.B2.A1.SetBytes(b[5*fr.Bytes : 6*fr.Bytes]) - return z + return z, nil } // MulByElement multiplies an element in E6 by an element in fr. diff --git a/internal/generator/field/template/extensions/e6_test.go.tmpl b/internal/generator/field/template/extensions/e6_test.go.tmpl index 9de29ff567..752fe6aabb 100644 --- a/internal/generator/field/template/extensions/e6_test.go.tmpl +++ b/internal/generator/field/template/extensions/e6_test.go.tmpl @@ -1,4 +1,5 @@ import ( + "github.com/stretchr/testify/require" "math/big" "testing" @@ -369,20 +370,17 @@ func TestE6MarshalSetBytesRoundTrip(t *testing.T) { } var y E6 - if !y.SetBytes(b).Equal(&x) { + _, err := y.SetBytes(b) + require.NoError(t, err) + if !y.Equal(&x) { t.Fatal("SetBytes(Marshal(x)) != x") } } for _, n := range []int{0, BytesE6 - 1, BytesE6 + 1} { - func() { - defer func() { - if recover() == nil { - t.Fatalf("SetBytes did not panic on %d bytes", n) - } - }() - var z E6 - z.SetBytes(make([]byte, n)) - }() + var z E6 + if _, err := z.SetBytes(make([]byte, n)); err == nil { + t.Fatalf("SetBytes did not fail on %d bytes", n) + } } } diff --git a/internal/generator/main.go b/internal/generator/main.go index d805c6b919..eb30909e29 100644 --- a/internal/generator/main.go +++ b/internal/generator/main.go @@ -92,15 +92,15 @@ func main() { )) // polynomial package (Polynomial, MultiLin, Pool, ...) over the - // degree-6 extension E6 - if f.Name == "koalabear" { + // requested extensions + for i, degree := range f.PolynomialExtensions { extInfo := fieldConfig.FieldDependency{ FieldPackagePath: "github.com/consensys/gnark-crypto/field/" + f.Name + "/extensions", FieldPackageName: "extensions", - ElementType: "extensions.E6", - ExtensionDegree: 6, + ElementType: fmt.Sprintf("extensions.E%d", degree), + ExtensionDegree: degree, } - assertNoError(polynomial.Generate(extInfo, filepath.Join(outputDir, "extensions", "polynomial"), true, gen)) + assertNoError(polynomial.Generate(extInfo, filepath.Join(outputDir, "extensions", "polynomial"), i == 0, true, gen)) } }(conf) } @@ -203,7 +203,7 @@ func main() { } assertNoError(mimc.Generate(conf, filepath.Join(curveDir, "fr", "mimc"), gen)) - assertNoError(polynomial.Generate(frInfo, filepath.Join(curveDir, "fr", "polynomial"), true, gen)) + assertNoError(polynomial.Generate(frInfo, filepath.Join(curveDir, "fr", "polynomial"), true, true, gen)) assertNoError(poseidon2.Generate(conf, filepath.Join(curveDir, "fr", "poseidon2"), gen)) assertNoError(hash_to_field.Generate(frInfo, filepath.Join(curveDir, "fr", "hash_to_field"), gen)) assertNoError(hash_to_field.Generate(fpInfo, filepath.Join(curveDir, "fp", "hash_to_field"), gen)) diff --git a/internal/generator/polynomial/generate.go b/internal/generator/polynomial/generate.go index 7739d6b4ea..bed38c1aee 100644 --- a/internal/generator/polynomial/generate.go +++ b/internal/generator/polynomial/generate.go @@ -2,6 +2,7 @@ package polynomial import ( "path/filepath" + "strings" "github.com/consensys/gnark-crypto/internal/generator/common" "github.com/consensys/gnark-crypto/internal/generator/field/config" @@ -10,18 +11,32 @@ import ( "github.com/consensys/gnark-crypto/internal/generator/polynomial/template" ) -func Generate(conf config.FieldDependency, baseDir string, generateTests bool, gen *common.Generator) error { - entries := []bavard.Entry{ - {File: filepath.Join(baseDir, "doc.go"), Templates: []string{"doc.go.tmpl"}}, - {File: filepath.Join(baseDir, "polynomial.go"), Templates: []string{"polynomial.go.tmpl"}}, - {File: filepath.Join(baseDir, "multilin.go"), Templates: []string{"multilin.go.tmpl"}}, - {File: filepath.Join(baseDir, "pool.go"), Templates: []string{"pool.go.tmpl"}}, +// Generate generates, in baseDir, the polynomial package over the field +// described by conf. For an extension, every free function and type is suffixed +// with the extension name (e.g. PolynomialE6) and the file names carry the +// lowercase name (e.g. polynomial_e6.go), so that several extensions can share +// a package. For a base field no suffix is added. doc.go is generated only when +// withDoc is set. +func Generate(conf config.FieldDependency, baseDir string, withDoc, generateTests bool, gen *common.Generator) error { + ext := "" + if name := conf.ExtensionName(); name != "" { + ext = "_" + strings.ToLower(name) } + var entries []bavard.Entry + if withDoc { + entries = append(entries, bavard.Entry{File: filepath.Join(baseDir, "doc.go"), Templates: []string{"doc.go.tmpl"}}) + } + entries = append(entries, + bavard.Entry{File: filepath.Join(baseDir, "polynomial"+ext+".go"), Templates: []string{"polynomial.go.tmpl"}}, + bavard.Entry{File: filepath.Join(baseDir, "multilin"+ext+".go"), Templates: []string{"multilin.go.tmpl"}}, + bavard.Entry{File: filepath.Join(baseDir, "pool"+ext+".go"), Templates: []string{"pool.go.tmpl"}}, + ) + if generateTests { entries = append(entries, - bavard.Entry{File: filepath.Join(baseDir, "polynomial_test.go"), Templates: []string{"polynomial.test.go.tmpl"}}, - bavard.Entry{File: filepath.Join(baseDir, "multilin_test.go"), Templates: []string{"multilin.test.go.tmpl"}}, + bavard.Entry{File: filepath.Join(baseDir, "polynomial"+ext+"_test.go"), Templates: []string{"polynomial.test.go.tmpl"}}, + bavard.Entry{File: filepath.Join(baseDir, "multilin"+ext+"_test.go"), Templates: []string{"multilin.test.go.tmpl"}}, ) } diff --git a/internal/generator/polynomial/template/multilin.go.tmpl b/internal/generator/polynomial/template/multilin.go.tmpl index 8cb631ddcf..50a08c10a5 100644 --- a/internal/generator/polynomial/template/multilin.go.tmpl +++ b/internal/generator/polynomial/template/multilin.go.tmpl @@ -4,14 +4,14 @@ import ( "github.com/consensys/gnark-crypto/utils" ) -// MultiLin tracks the values of a (dense i.e. not sparse) multilinear polynomial +// MultiLin{{.ExtensionName}} tracks the values of a (dense i.e. not sparse) multilinear polynomial // The variables are X₁ through Xₙ where n = log(len(.)) // .[∑ᵢ 2ⁱ⁻¹ bₙ₋ᵢ] = the polynomial evaluated at (b₁, b₂, ..., bₙ) // It is understood that any hypercube evaluation can be extrapolated to a multilinear polynomial -type MultiLin []{{.ElementType}} +type MultiLin{{.ExtensionName}} []{{.ElementType}} // Fold is partial evaluation function k[X₁, X₂, ..., Xₙ] → k[X₂, ..., Xₙ] by setting X₁=r -func (m *MultiLin) Fold(r {{.ElementType}}) { +func (m *MultiLin{{.ExtensionName}}) Fold(r {{.ElementType}}) { mid := len(*m) / 2 bottom, top := (*m)[:mid], (*m)[mid:] @@ -32,7 +32,7 @@ func (m *MultiLin) Fold(r {{.ElementType}}) { *m = (*m)[:mid] } -func (m *MultiLin) FoldParallel(r {{.ElementType}}) utils.Task { +func (m *MultiLin{{.ExtensionName}}) FoldParallel(r {{.ElementType}}) utils.Task { mid := len(*m) / 2 bottom, top := (*m)[:mid], (*m)[mid:] @@ -49,7 +49,7 @@ func (m *MultiLin) FoldParallel(r {{.ElementType}}) utils.Task { } } -func (m MultiLin) Sum() {{.ElementType}} { +func (m MultiLin{{.ExtensionName}}) Sum() {{.ElementType}} { s := m[0] for i := 1; i < len(m); i++ { s.Add(&s, &m[i]) @@ -57,7 +57,7 @@ func (m MultiLin) Sum() {{.ElementType}} { return s } -func _clone(m MultiLin, p *Pool) MultiLin { +func _clone{{.ExtensionName}}(m MultiLin{{.ExtensionName}}, p *Pool{{.ExtensionName}}) MultiLin{{.ExtensionName}} { if p == nil { return m.Clone() } else { @@ -65,7 +65,7 @@ func _clone(m MultiLin, p *Pool) MultiLin { } } -func _dump(m MultiLin, p *Pool) { +func _dump{{.ExtensionName}}(m MultiLin{{.ExtensionName}}, p *Pool{{.ExtensionName}}) { if p != nil { p.Dump(m) } @@ -73,9 +73,9 @@ func _dump(m MultiLin, p *Pool) { // Evaluate extrapolate the value of the multilinear polynomial corresponding to m // on the given coordinates -func (m MultiLin) Evaluate(coordinates []{{.ElementType}}, p *Pool) {{.ElementType}} { +func (m MultiLin{{.ExtensionName}}) Evaluate(coordinates []{{.ElementType}}, p *Pool{{.ExtensionName}}) {{.ElementType}} { // Folding is a mutating operation - bkCopy := _clone(m, p) + bkCopy := _clone{{.ExtensionName}}(m, p) // Evaluate step by step through repeated folding (i.e. evaluation at the first remaining variable) for _, r := range coordinates { @@ -84,7 +84,7 @@ func (m MultiLin) Evaluate(coordinates []{{.ElementType}}, p *Pool) {{.ElementTy result := bkCopy[0] - _dump(bkCopy, p) + _dump{{.ExtensionName}}(bkCopy, p) return result } @@ -92,14 +92,14 @@ func (m MultiLin) Evaluate(coordinates []{{.ElementType}}, p *Pool) {{.ElementTy // Both multilinear interpolation and sumcheck require folding an underlying // array, but folding changes the array. To do both one requires a deep copy // of the bookkeeping table. -func (m MultiLin) Clone() MultiLin { - res := make(MultiLin, len(m)) +func (m MultiLin{{.ExtensionName}}) Clone() MultiLin{{.ExtensionName}} { + res := make(MultiLin{{.ExtensionName}}, len(m)) copy(res, m) return res } // Add two bookKeepingTables -func (m *MultiLin) Add(left, right MultiLin) { +func (m *MultiLin{{.ExtensionName}}) Add(left, right MultiLin{{.ExtensionName}}) { size := len(left) // Check that left and right have the same size if len(right) != size || len(*m) != size{ @@ -113,7 +113,7 @@ func (m *MultiLin) Add(left, right MultiLin) { } -// EvalEq computes Eq(q₁, ... , qₙ, h₁, ... , hₙ) = Π₁ⁿ Eq(qᵢ, hᵢ) +// EvalEq{{.ExtensionName}} computes Eq(q₁, ... , qₙ, h₁, ... , hₙ) = Π₁ⁿ Eq(qᵢ, hᵢ) // where Eq(x,y) = xy + (1-x)(1-y) = 1 - x - y + xy + xy interpolates // _________________ // | | | @@ -126,7 +126,7 @@ func (m *MultiLin) Add(left, right MultiLin) { // x // In other words the polynomial evaluated here is the multilinear extrapolation of // one that evaluates to q' == h' for vectors q', h' of binary values -func EvalEq(q, h []{{.ElementType}}) {{.ElementType}} { +func EvalEq{{.ExtensionName}}(q, h []{{.ElementType}}) {{.ElementType}} { var res, nxt, one, sum {{.ElementType}} one.SetOne() for i := range len(q) { @@ -146,7 +146,7 @@ func EvalEq(q, h []{{.ElementType}}) {{.ElementType}} { } // Eq sets m to the representation of the polynomial Eq(q₁, ..., qₙ, *, ..., *) × m[0] -func (m *MultiLin) Eq(q []{{.ElementType}}) { +func (m *MultiLin{{.ExtensionName}}) Eq(q []{{.ElementType}}) { n := len(q) if len(*m) != 1 << n { @@ -165,6 +165,6 @@ func (m *MultiLin) Eq(q []{{.ElementType}}) { } } -func (m MultiLin) NumVars() int { +func (m MultiLin{{.ExtensionName}}) NumVars() int { return bits.TrailingZeros(uint(len(m))) } \ No newline at end of file diff --git a/internal/generator/polynomial/template/multilin.test.go.tmpl b/internal/generator/polynomial/template/multilin.test.go.tmpl index c71bf080f3..878b05aa90 100644 --- a/internal/generator/polynomial/template/multilin.test.go.tmpl +++ b/internal/generator/polynomial/template/multilin.test.go.tmpl @@ -5,14 +5,14 @@ import ( ) //TODO: Property based tests? -func TestFoldBilinear(t *testing.T) { +func TestFoldBilinear{{.ExtensionName}}(t *testing.T) { for range 100 { // f = c₀ + c₁ X₁ + c₂ X₂ + c₃ X₁ X₂ var coefficients [4]{{.ElementType}} {{- if .ExtensionDegree}} - {{.FieldPackageName}}.VectorE{{.ExtensionDegree}}(coefficients[:]).MustSetRandom() + {{.FieldPackageName}}.Vector{{.ExtensionName}}(coefficients[:]).MustSetRandom() {{- else}} {{.FieldPackageName}}.Vector(coefficients[:]).MustSetRandom() {{- end}} @@ -21,7 +21,7 @@ func TestFoldBilinear(t *testing.T) { r.MustSetRandom() // interpolate at {0,1}²: - m := make(MultiLin, 4) + m := make(MultiLin{{.ExtensionName}}, 4) m[0] = coefficients[0] m[1].Add(&coefficients[0], &coefficients[2]) m[2].Add(&coefficients[0], &coefficients[1]) @@ -50,12 +50,12 @@ func TestFoldBilinear(t *testing.T) { // TODO: Benchmark folding? Algorithms is pretty straightforward; unless we want to measure how well memory management is working -func TestFoldedEqTable(t *testing.T) { +func TestFoldedEqTable{{.ExtensionName}}(t *testing.T) { q := make([]{{.ElementType}}, 2) q[0].SetInt64(2) q[1].SetInt64(3) - m := make(MultiLin, 4) + m := make(MultiLin{{.ExtensionName}}, 4) m[0].SetOne() m.Eq(q) @@ -68,7 +68,7 @@ func TestFoldedEqTable(t *testing.T) { for p0 := range 2 { p[1].SetZero() for p1 := range 2 { - eq[p0*2+p1] = EvalEq(q, p) + eq[p0*2+p1] = EvalEq{{.ExtensionName}}(q, p) p[1].Add(&p[1], &one) } p[0].Add(&p[0], &one) diff --git a/internal/generator/polynomial/template/polynomial.go.tmpl b/internal/generator/polynomial/template/polynomial.go.tmpl index 0efdd422c6..a604d5be12 100644 --- a/internal/generator/polynomial/template/polynomial.go.tmpl +++ b/internal/generator/polynomial/template/polynomial.go.tmpl @@ -6,17 +6,17 @@ import ( "sync" ) -// Polynomial represented by coefficients in the field. -type Polynomial []{{.ElementType}} +// Polynomial{{.ExtensionName}} represented by coefficients in the field. +type Polynomial{{.ExtensionName}} []{{.ElementType}} // Degree returns the degree of the polynomial, which is the length of Data. -func (p *Polynomial) Degree() uint64 { +func (p *Polynomial{{.ExtensionName}}) Degree() uint64 { return uint64(len(*p) - 1) } // Eval evaluates p at v // returns a {{.ElementType}} -func (p *Polynomial) Eval(v *{{.ElementType}}) {{.ElementType}} { +func (p *Polynomial{{.ExtensionName}}) Eval(v *{{.ElementType}}) {{.ElementType}} { res := (*p)[len(*p) - 1] for i := len(*p) - 2; i >= 0; i-- { @@ -28,14 +28,14 @@ func (p *Polynomial) Eval(v *{{.ElementType}}) {{.ElementType}} { } // Clone returns a copy of the polynomial -func (p *Polynomial) Clone() Polynomial { - _p := make(Polynomial, len(*p)) +func (p *Polynomial{{.ExtensionName}}) Clone() Polynomial{{.ExtensionName}} { + _p := make(Polynomial{{.ExtensionName}}, len(*p)) copy(_p, *p) return _p } // Set to another polynomial -func (p *Polynomial) Set(p1 Polynomial) { +func (p *Polynomial{{.ExtensionName}}) Set(p1 Polynomial{{.ExtensionName}}) { if len(*p) != len(p1) { *p = p1.Clone() return @@ -47,30 +47,30 @@ func (p *Polynomial) Set(p1 Polynomial) { } // AddConstantInPlace adds a constant to the polynomial, modifying p -func (p *Polynomial) AddConstantInPlace(c *{{.ElementType}}) { +func (p *Polynomial{{.ExtensionName}}) AddConstantInPlace(c *{{.ElementType}}) { for i := range len(*p) { (*p)[i].Add(&(*p)[i], c) } } // SubConstantInPlace subs a constant to the polynomial, modifying p -func (p *Polynomial) SubConstantInPlace(c *{{.ElementType}}) { +func (p *Polynomial{{.ExtensionName}}) SubConstantInPlace(c *{{.ElementType}}) { for i := range len(*p) { (*p)[i].Sub(&(*p)[i], c) } } // ScaleInPlace multiplies p by v, modifying p -func (p *Polynomial) ScaleInPlace(c *{{.ElementType}}) { +func (p *Polynomial{{.ExtensionName}}) ScaleInPlace(c *{{.ElementType}}) { for i := range len(*p) { (*p)[i].Mul(&(*p)[i], c) } } // Scale multiplies p0 by v, storing the result in p -func (p *Polynomial) Scale(c *{{.ElementType}}, p0 Polynomial) { +func (p *Polynomial{{.ExtensionName}}) Scale(c *{{.ElementType}}, p0 Polynomial{{.ExtensionName}}) { if len(*p) != len(p0) { - *p = make(Polynomial, len(p0)) + *p = make(Polynomial{{.ExtensionName}}, len(p0)) } for i := range len(p0) { (*p)[i].Mul(c, &p0[i]) @@ -79,7 +79,7 @@ func (p *Polynomial) Scale(c *{{.ElementType}}, p0 Polynomial) { // Add adds p1 to p2 // This function allocates a new slice unless p == p1 or p == p2 -func (p *Polynomial) Add(p1, p2 Polynomial) *Polynomial { +func (p *Polynomial{{.ExtensionName}}) Add(p1, p2 Polynomial{{.ExtensionName}}) *Polynomial{{.ExtensionName}} { bigger := p1 smaller := p2 @@ -102,7 +102,7 @@ func (p *Polynomial) Add(p1, p2 Polynomial) *Polynomial { return p } - res := make(Polynomial, len(bigger)) + res := make(Polynomial{{.ExtensionName}}, len(bigger)) copy(res, bigger) for i := range len(smaller) { res[i].Add(&res[i], &smaller[i]) @@ -113,7 +113,7 @@ func (p *Polynomial) Add(p1, p2 Polynomial) *Polynomial { // Sub subtracts p2 from p1 // TODO make interface more consistent with Add -func (p *Polynomial) Sub(p1, p2 Polynomial) *Polynomial { +func (p *Polynomial{{.ExtensionName}}) Sub(p1, p2 Polynomial{{.ExtensionName}}) *Polynomial{{.ExtensionName}} { if len(p1) != len(p2) || len(p2) != len(*p) { return nil } @@ -124,7 +124,7 @@ func (p *Polynomial) Sub(p1, p2 Polynomial) *Polynomial { } // Equal checks equality between two polynomials -func (p *Polynomial) Equal(p1 Polynomial) bool { +func (p *Polynomial{{.ExtensionName}}) Equal(p1 Polynomial{{.ExtensionName}}) bool { if (*p == nil) != (p1 == nil) { return false } @@ -142,14 +142,14 @@ func (p *Polynomial) Equal(p1 Polynomial) bool { return true } -func (p Polynomial) SetZero() { +func (p Polynomial{{.ExtensionName}}) SetZero() { for i := range len(p) { p[i].SetZero() } } {{if not .ExtensionDegree}} -func (p Polynomial) Text(base int) string { +func (p Polynomial{{.ExtensionName}}) Text(base int) string { var builder strings.Builder @@ -204,20 +204,20 @@ func (p Polynomial) Text(base int) string { } {{end}} -// InterpolateOnRange maps vector v to polynomial f +// InterpolateOnRange{{.ExtensionName}} maps vector v to polynomial f // such that f(i) = v[i] for 0 ≤ i < len(v). // len(f) = len(v) and deg(f) ≤ len(v) - 1 -func InterpolateOnRange(v []{{.ElementType}}) Polynomial { +func InterpolateOnRange{{.ExtensionName}}(v []{{.ElementType}}) Polynomial{{.ExtensionName}} { nEvals := uint8(len(v)) if int(nEvals) != len(v) { panic("interpolation method too inefficient for nEvals > 255") } - lagrange := getLagrangeBasis(nEvals) + lagrange := getLagrangeBasis{{.ExtensionName}}(nEvals) - var res Polynomial + var res Polynomial{{.ExtensionName}} res.Scale(&v[0], lagrange[0]) - temp := make(Polynomial, nEvals) + temp := make(Polynomial{{.ExtensionName}}, nEvals) for i := uint8(1); i < nEvals; i++ { temp.Scale(&v[i], lagrange[i]) @@ -227,39 +227,39 @@ func InterpolateOnRange(v []{{.ElementType}}) Polynomial { return res } -// lagrange bases used by InterpolateOnRange -var lagrangeBasis sync.Map +// lagrange bases used by InterpolateOnRange{{.ExtensionName}} +var lagrangeBasis{{.ExtensionName}} sync.Map -func getLagrangeBasis(domainSize uint8) []Polynomial { - if res, ok := lagrangeBasis.Load(domainSize); ok { - return res.([]Polynomial) +func getLagrangeBasis{{.ExtensionName}}(domainSize uint8) []Polynomial{{.ExtensionName}} { + if res, ok := lagrangeBasis{{.ExtensionName}}.Load(domainSize); ok { + return res.([]Polynomial{{.ExtensionName}}) } // not found. compute - var res []Polynomial + var res []Polynomial{{.ExtensionName}} if domainSize >= 2 { - res = computeLagrangeBasis(domainSize) + res = computeLagrangeBasis{{.ExtensionName}}(domainSize) } else if domainSize == 1 { - res = []Polynomial{make(Polynomial, 1)} + res = []Polynomial{{.ExtensionName}}{make(Polynomial{{.ExtensionName}}, 1)} res[0][0].SetOne() } - lagrangeBasis.Store(domainSize, res) + lagrangeBasis{{.ExtensionName}}.Store(domainSize, res) return res } -// computeLagrangeBasis precomputes in explicit coefficient form for each 0 ≤ l < domainSize the polynomial +// computeLagrangeBasis{{.ExtensionName}} precomputes in explicit coefficient form for each 0 ≤ l < domainSize the polynomial // pₗ := X (X-1) ... (X-l-1) (X-l+1) ... (X - domainSize + 1) / ( l (l-1) ... 2 (-1) ... (l - domainSize +1) ) // Note that pₗ(l) = 1 and pₗ(n) = 0 if 0 ≤ l < domainSize, n ≠ l -func computeLagrangeBasis(domainSize uint8) []Polynomial { +func computeLagrangeBasis{{.ExtensionName}}(domainSize uint8) []Polynomial{{.ExtensionName}} { constTerms := make([]{{.ElementType}}, domainSize) for i := range domainSize { constTerms[i].SetInt64(-int64(i)) } - res := make([]Polynomial, domainSize) - multScratch := make(Polynomial, domainSize-1) + res := make([]Polynomial{{.ExtensionName}}, domainSize) + multScratch := make(Polynomial{{.ExtensionName}}, domainSize-1) // compute pₗ for l := range domainSize { @@ -271,7 +271,7 @@ func computeLagrangeBasis(domainSize uint8) []Polynomial { continue } if d == 0 { - res[l] = make(Polynomial, domainSize) + res[l] = make(Polynomial{{.ExtensionName}}, domainSize) res[l][domainSize-2] = constTerms[i] res[l][domainSize-1].SetOne() } else { @@ -296,7 +296,7 @@ func computeLagrangeBasis(domainSize uint8) []Polynomial { constTerms[l] = res[l].Eval(&constTerms[l]) } {{- if .ExtensionDegree}} - constTerms = {{.FieldPackageName}}.BatchInvertE{{.ExtensionDegree}}(constTerms) + constTerms = {{.FieldPackageName}}.BatchInvert{{.ExtensionName}}(constTerms) {{- else}} constTerms = {{.FieldPackageName}}.BatchInvert(constTerms) {{- end}} diff --git a/internal/generator/polynomial/template/polynomial.test.go.tmpl b/internal/generator/polynomial/template/polynomial.test.go.tmpl index da0af09257..a7c7bbb1e5 100644 --- a/internal/generator/polynomial/template/polynomial.test.go.tmpl +++ b/internal/generator/polynomial/template/polynomial.test.go.tmpl @@ -8,10 +8,10 @@ import ( "testing" ) -func TestPolynomialEval(t *testing.T) { +func TestPolynomialEval{{.ExtensionName}}(t *testing.T) { // build polynomial - f := make(Polynomial, 20) + f := make(Polynomial{{.ExtensionName}}, 20) for i := range 20 { f[i].SetOne() } @@ -39,10 +39,10 @@ func TestPolynomialEval(t *testing.T) { } } -func TestPolynomialAddConstantInPlace(t *testing.T) { +func TestPolynomialAddConstantInPlace{{.ExtensionName}}(t *testing.T) { // build polynomial - f := make(Polynomial, 20) + f := make(Polynomial{{.ExtensionName}}, 20) for i := range 20 { f[i].SetOne() } @@ -65,10 +65,10 @@ func TestPolynomialAddConstantInPlace(t *testing.T) { } } -func TestPolynomialSubConstantInPlace(t *testing.T) { +func TestPolynomialSubConstantInPlace{{.ExtensionName}}(t *testing.T) { // build polynomial - f := make(Polynomial, 20) + f := make(Polynomial{{.ExtensionName}}, 20) for i := range 20 { f[i].SetOne() } @@ -91,10 +91,10 @@ func TestPolynomialSubConstantInPlace(t *testing.T) { } } -func TestPolynomialScaleInPlace(t *testing.T) { +func TestPolynomialScaleInPlace{{.ExtensionName}}(t *testing.T) { // build polynomial - f := make(Polynomial, 20) + f := make(Polynomial{{.ExtensionName}}, 20) for i := range 20 { f[i].SetOne() } @@ -115,17 +115,17 @@ func TestPolynomialScaleInPlace(t *testing.T) { } -func TestPolynomialAdd(t *testing.T) { +func TestPolynomialAdd{{.ExtensionName}}(t *testing.T) { // build unbalanced polynomials - f1 := make(Polynomial, 20) - f1Backup := make(Polynomial, 20) + f1 := make(Polynomial{{.ExtensionName}}, 20) + f1Backup := make(Polynomial{{.ExtensionName}}, 20) for i := range 20 { f1[i].SetOne() f1Backup[i].SetOne() } - f2 := make(Polynomial, 10) - f2Backup := make(Polynomial, 10) + f2 := make(Polynomial{{.ExtensionName}}, 10) + f2Backup := make(Polynomial{{.ExtensionName}}, 10) for i := range 10 { f2[i].SetOne() f2Backup[i].SetOne() @@ -135,7 +135,7 @@ func TestPolynomialAdd(t *testing.T) { var one, two {{.ElementType}} one.SetOne() two.Double(&one) - expectedSum := make(Polynomial, 20) + expectedSum := make(Polynomial{{.ExtensionName}}, 20) for i := range 10 { expectedSum[i].Set(&two) } @@ -144,7 +144,7 @@ func TestPolynomialAdd(t *testing.T) { } // caller is empty - var g Polynomial + var g Polynomial{{.ExtensionName}} g.Add(f1, f2) if !g.Equal(expectedSum) { t.Fatal("add polynomials fails") @@ -193,21 +193,21 @@ func TestPolynomialAdd(t *testing.T) { } {{if not .ExtensionDegree}} -func TestPolynomialText(t *testing.T) { +func TestPolynomialText{{.ExtensionName}}(t *testing.T) { var one, negTwo {{.ElementType}} one.SetOne() negTwo.SetInt64(-2) - p := Polynomial{one, negTwo, one} + p := Polynomial{{.ExtensionName}}{one, negTwo, one} assert.Equal(t, "X² - 2X + 1", p.Text(10)) } {{end}} -func TestPrecomputeLagrange(t *testing.T) { +func TestPrecomputeLagrange{{.ExtensionName}}(t *testing.T) { testForDomainSize := func(domainSize uint8) bool { - polys := computeLagrangeBasis(domainSize) + polys := computeLagrangeBasis{{.ExtensionName}}(domainSize) for l := range domainSize { for i := range domainSize { @@ -245,9 +245,9 @@ func TestPrecomputeLagrange(t *testing.T) { properties.TestingRun(t, gopter.ConsoleReporter(false)) } -func TestLagrangeCache(t *testing.T) { +func TestLagrangeCache{{.ExtensionName}}(t *testing.T) { for _, i := range []int{5, 2, 8, 4, 6, 3, 0} { - b := getLagrangeBasis(uint8(i)) - assert.Equal(t, b, getLagrangeBasis(uint8(i))) // second call must yield the same result + b := getLagrangeBasis{{.ExtensionName}}(uint8(i)) + assert.Equal(t, b, getLagrangeBasis{{.ExtensionName}}(uint8(i))) // second call must yield the same result } } \ No newline at end of file diff --git a/internal/generator/polynomial/template/pool.go.tmpl b/internal/generator/polynomial/template/pool.go.tmpl index a3b3a9487c..9314356e40 100644 --- a/internal/generator/polynomial/template/pool.go.tmpl +++ b/internal/generator/polynomial/template/pool.go.tmpl @@ -12,38 +12,38 @@ import ( // WARNING: This is not thread safe TODO: Make sure that is not a problem // TODO: There is a lot of "unsafe" memory management here and needs to be vetted thoroughly -type sizedPool struct { +type sizedPool{{.ExtensionName}} struct { maxN int pool sync.Pool - stats poolStats + stats poolStats{{.ExtensionName}} } -type inUseData struct { +type inUseData{{.ExtensionName}} struct { allocatedFor []uintptr - pool *sizedPool + pool *sizedPool{{.ExtensionName}} } -type Pool struct { +type Pool{{.ExtensionName}} struct { //lock sync.Mutex inUse sync.Map - subPools []sizedPool + subPools []sizedPool{{.ExtensionName}} } -func (p *sizedPool) get(n int) *{{.ElementType}} { +func (p *sizedPool{{.ExtensionName}}) get(n int) *{{.ElementType}} { p.stats.make(n) return p.pool.Get().(*{{.ElementType}}) } -func (p *sizedPool) put(ptr *{{.ElementType}}) { +func (p *sizedPool{{.ExtensionName}}) put(ptr *{{.ElementType}}) { p.stats.dump() p.pool.Put(ptr) } -func NewPool(maxN ...int) (pool Pool) { +func NewPool{{.ExtensionName}}(maxN ...int) (pool Pool{{.ExtensionName}}) { sort.Ints(maxN) - pool = Pool{ - subPools: make([]sizedPool, len(maxN)), + pool = Pool{{.ExtensionName}}{ + subPools: make([]sizedPool{{.ExtensionName}}, len(maxN)), } for i := range pool.subPools { @@ -52,14 +52,14 @@ func NewPool(maxN ...int) (pool Pool) { subPool.pool = sync.Pool{ New: func() any { subPool.stats.Allocated++ - return getDataPointer(make([]{{.ElementType}}, 0, subPool.maxN)) + return getDataPointer{{.ExtensionName}}(make([]{{.ElementType}}, 0, subPool.maxN)) }, } } return } -func (p *Pool) findCorrespondingPool(n int) *sizedPool { +func (p *Pool{{.ExtensionName}}) findCorrespondingPool(n int) *sizedPool{{.ExtensionName}} { poolI := 0 for poolI < len(p.subPools) && n > p.subPools[poolI].maxN { poolI++ @@ -67,7 +67,7 @@ func (p *Pool) findCorrespondingPool(n int) *sizedPool { return &p.subPools[poolI] // out of bounds error here would mean that n is too large } -func (p *Pool) Make(n int) []{{.ElementType}} { +func (p *Pool{{.ExtensionName}}) Make(n int) []{{.ElementType}} { pool := p.findCorrespondingPool(n) ptr := pool.get(n) p.addInUse(ptr, pool) @@ -75,52 +75,52 @@ func (p *Pool) Make(n int) []{{.ElementType}} { } // Dump dumps a set of polynomials into the pool -func (p *Pool) Dump(slices ...[]{{.ElementType}}) { +func (p *Pool{{.ExtensionName}}) Dump(slices ...[]{{.ElementType}}) { for _, slice := range slices { - ptr := getDataPointer(slice) + ptr := getDataPointer{{.ExtensionName}}(slice) if metadata, ok := p.inUse.Load(ptr); ok { p.inUse.Delete(ptr) - metadata.(inUseData).pool.put(ptr) + metadata.(inUseData{{.ExtensionName}}).pool.put(ptr) } else { panic("attempting to dump a slice not created by the pool") } } } -func (p *Pool) addInUse(ptr *{{.ElementType}}, pool *sizedPool) { +func (p *Pool{{.ExtensionName}}) addInUse(ptr *{{.ElementType}}, pool *sizedPool{{.ExtensionName}}) { pcs := make([]uintptr, 2) n := runtime.Callers(3, pcs) if prevPcs, ok := p.inUse.Load(ptr); ok { // TODO: remove if unnecessary for security - panic(fmt.Errorf("re-allocated non-dumped slice, previously allocated at %v", runtime.CallersFrames(prevPcs.(inUseData).allocatedFor))) + panic(fmt.Errorf("re-allocated non-dumped slice, previously allocated at %v", runtime.CallersFrames(prevPcs.(inUseData{{.ExtensionName}}).allocatedFor))) } - p.inUse.Store(ptr, inUseData{ + p.inUse.Store(ptr, inUseData{{.ExtensionName}}{ allocatedFor: pcs[:n], pool: pool, }) } -func printFrame(frame runtime.Frame) { +func printFrame{{.ExtensionName}}(frame runtime.Frame) { fmt.Printf("\t%s line %d, function %s\n", frame.File, frame.Line, frame.Function) } -func (p *Pool) printInUse() { +func (p *Pool{{.ExtensionName}}) printInUse() { fmt.Println("slices never dumped allocated at:") p.inUse.Range(func(_, pcs any) bool { fmt.Println("-------------------------") var frame runtime.Frame - frames := runtime.CallersFrames(pcs.(inUseData).allocatedFor) + frames := runtime.CallersFrames(pcs.(inUseData{{.ExtensionName}}).allocatedFor) more := true for more { frame, more = frames.Next() - printFrame(frame) + printFrame{{.ExtensionName}}(frame) } return true }) } -type poolStats struct { +type poolStats{{.ExtensionName}} struct { Used int Allocated int ReuseRate float64 @@ -129,12 +129,12 @@ type poolStats struct { SmallestNUsed int } -type poolsStats struct { - SubPools []poolStats +type poolsStats{{.ExtensionName}} struct { + SubPools []poolStats{{.ExtensionName}} InUse int } -func (s *poolStats) make(n int) { +func (s *poolStats{{.ExtensionName}}) make(n int) { s.Used++ s.InUse++ if n > s.GreatestNUsed { @@ -145,21 +145,21 @@ func (s *poolStats) make(n int) { } } -func (s *poolStats) dump() { +func (s *poolStats{{.ExtensionName}}) dump() { s.InUse-- } -func (s *poolStats) finalize() { +func (s *poolStats{{.ExtensionName}}) finalize() { s.ReuseRate = float64(s.Used) / float64(s.Allocated) } -func getDataPointer(slice []{{.ElementType}}) *{{.ElementType}} { +func getDataPointer{{.ExtensionName}}(slice []{{.ElementType}}) *{{.ElementType}} { return (*{{.ElementType}})(unsafe.SliceData(slice)) } -func (p *Pool) PrintPoolStats() { +func (p *Pool{{.ExtensionName}}) PrintPoolStats() { InUse := 0 - subStats := make([]poolStats, len(p.subPools)) + subStats := make([]poolStats{{.ExtensionName}}, len(p.subPools)) for i := range p.subPools { subPool := &p.subPools[i] subPool.stats.finalize() @@ -167,7 +167,7 @@ func (p *Pool) PrintPoolStats() { InUse += subPool.stats.InUse } - stats := poolsStats{ + stats := poolsStats{{.ExtensionName}}{ SubPools: subStats, InUse: InUse, } @@ -176,7 +176,7 @@ func (p *Pool) PrintPoolStats() { p.printInUse() } -func (p *Pool) Clone(slice []{{.ElementType}}) []{{.ElementType}} { +func (p *Pool{{.ExtensionName}}) Clone(slice []{{.ElementType}}) []{{.ElementType}} { res := p.Make(len(slice)) copy(res, slice) return res From bb49c2d81876aeb4bdfc57dc3150590e974a6022 Mon Sep 17 00:00:00 2001 From: Arya Tabaie Date: Mon, 5 Oct 2026 13:44:55 -0500 Subject: [PATCH 08/18] fix: cursor findings Signed-off-by: Arya Tabaie --- field/babybear/extensions/e2.go | 14 ++ field/babybear/extensions/e2_test.go | 28 +++ field/babybear/extensions/e4.go | 14 ++ field/babybear/extensions/e4_test.go | 160 +++++++++++------- field/babybear/extensions/vector.go | 87 +++++----- field/babybear/fft/fftext.go | 18 +- field/goldilocks/extensions/e2.go | 14 ++ field/goldilocks/extensions/e2_test.go | 28 +++ field/koalabear/extensions/e2.go | 14 ++ field/koalabear/extensions/e2_test.go | 28 +++ field/koalabear/extensions/e4.go | 14 ++ field/koalabear/extensions/e4_test.go | 160 +++++++++++------- field/koalabear/extensions/vector.go | 87 +++++----- field/koalabear/fft/fftext.go | 18 +- field/koalabear/vortex/batch_poly.go | 8 +- field/koalabear/vortex/prover.go | 2 +- field/mamabear/extensions/e3.go | 21 +++ field/mamabear/extensions/e3_test.go | 78 +++++++-- field/mamabear/extensions/e3_vector.go | 27 +-- .../extensions/e3_vector_internal_test.go | 6 +- field/mamabear/fft/fftext.go | 18 +- field/mamabear/vortex/batch_poly.go | 8 +- field/mamabear/vortex/prover.go | 2 +- .../field/template/extensions/e2.go.tmpl | 14 ++ .../field/template/extensions/e2_test.go.tmpl | 28 +++ .../field/template/extensions/e3.go.tmpl | 21 +++ .../field/template/extensions/e3_test.go.tmpl | 78 +++++++-- .../template/extensions/e3vector.go.tmpl | 27 +-- .../template/extensions/e3vector_test.go.tmpl | 6 +- .../field/template/extensions/e4.go.tmpl | 14 ++ .../field/template/extensions/e4_test.go.tmpl | 160 +++++++++++------- .../field/template/extensions/vector.go.tmpl | 87 +++++----- .../field/template/fft/fftext.go.tmpl | 18 +- 33 files changed, 890 insertions(+), 417 deletions(-) diff --git a/field/babybear/extensions/e2.go b/field/babybear/extensions/e2.go index 033a3e8ef0..8672c2de42 100644 --- a/field/babybear/extensions/e2.go +++ b/field/babybear/extensions/e2.go @@ -75,6 +75,20 @@ func (z *E2) SetOne() *E2 { return z } +// SetInt64 sets z to v (embedded via the unique ring homomorphism ℤ→𝔽) and returns z +func (z *E2) SetInt64(v int64) *E2 { + *z = E2{} + z.A0.SetInt64(v) + return z +} + +// SetUint64 sets z to v (embedded via the unique ring homomorphism ℤ→𝔽) and returns z +func (z *E2) SetUint64(v uint64) *E2 { + *z = E2{} + z.A0.SetUint64(v) + return z +} + // SetBigInt sets z to v (embedded via the unique ring homomorphism ℤ→𝔽, reduced modulo the base field order) and returns z func (z *E2) SetBigInt(v *big.Int) *E2 { *z = E2{} diff --git a/field/babybear/extensions/e2_test.go b/field/babybear/extensions/e2_test.go index 05e44f1ee6..da9cb653ef 100644 --- a/field/babybear/extensions/e2_test.go +++ b/field/babybear/extensions/e2_test.go @@ -589,3 +589,31 @@ func TestE2MarshalSetBytesRoundTrip(t *testing.T) { } } } + +func TestE2SetInt64SetUint64(t *testing.T) { + for _, v := range []int64{0, 1, -1, 7, -12345, 1 << 40, -(1 << 40)} { + var z E2 + z.SetInt64(v) + require.Truef(t, z.A1.IsZero(), "coordinate A1 is non-zero for %d", v) + var want fr.Element + want.SetInt64(v) + require.Truef(t, z.A0.Equal(&want), "SetInt64(%d)", v) + + // SetInt64 and SetBigInt agree + var zb E2 + zb.SetBigInt(big.NewInt(v)) + require.Truef(t, z.Equal(&zb), "SetInt64(%d) != SetBigInt(%d)", v, v) + } + for _, v := range []uint64{0, 1, 7, 12345, 1 << 40, 1<<64 - 1} { + var z E2 + z.SetUint64(v) + require.Truef(t, z.A1.IsZero(), "coordinate A1 is non-zero for %d", v) + var want fr.Element + want.SetUint64(v) + require.Truef(t, z.A0.Equal(&want), "SetUint64(%d)", v) + + var zb E2 + zb.SetBigInt(new(big.Int).SetUint64(v)) + require.Truef(t, z.Equal(&zb), "SetUint64(%d) != SetBigInt(%d)", v, v) + } +} diff --git a/field/babybear/extensions/e4.go b/field/babybear/extensions/e4.go index f6aba77ab2..811a714985 100644 --- a/field/babybear/extensions/e4.go +++ b/field/babybear/extensions/e4.go @@ -86,6 +86,20 @@ func (z *E4) SetOne() *E4 { return z } +// SetInt64 sets z to v (embedded via the unique ring homomorphism ℤ→𝔽) and returns z +func (z *E4) SetInt64(v int64) *E4 { + *z = E4{} + z.B0.A0.SetInt64(v) + return z +} + +// SetUint64 sets z to v (embedded via the unique ring homomorphism ℤ→𝔽) and returns z +func (z *E4) SetUint64(v uint64) *E4 { + *z = E4{} + z.B0.A0.SetUint64(v) + return z +} + // SetBigInt sets z to v (embedded via the unique ring homomorphism ℤ→𝔽, reduced modulo the base field order) and returns z func (z *E4) SetBigInt(v *big.Int) *E4 { *z = E4{} diff --git a/field/babybear/extensions/e4_test.go b/field/babybear/extensions/e4_test.go index 4becccf8d7..687f317e48 100644 --- a/field/babybear/extensions/e4_test.go +++ b/field/babybear/extensions/e4_test.go @@ -293,8 +293,8 @@ func TestVectorOps(t *testing.T) { } properties := gopter.NewProperties(parameters) - addVector := func(a, b Vector) bool { - c := make(Vector, len(a)) + addVector := func(a, b VectorE4) bool { + c := make(VectorE4, len(a)) c.Add(a, b) for i := range len(a) { @@ -307,8 +307,8 @@ func TestVectorOps(t *testing.T) { return true } - subVector := func(a, b Vector) bool { - c := make(Vector, len(a)) + subVector := func(a, b VectorE4) bool { + c := make(VectorE4, len(a)) c.Sub(a, b) for i := range len(a) { @@ -321,8 +321,8 @@ func TestVectorOps(t *testing.T) { return true } - scalarMulVector := func(a Vector, b E4) bool { - c := make(Vector, len(a)) + scalarMulVector := func(a VectorE4, b E4) bool { + c := make(VectorE4, len(a)) c.ScalarMul(a, &b) for i := range len(a) { @@ -335,7 +335,7 @@ func TestVectorOps(t *testing.T) { return true } - sumVector := func(a Vector) bool { + sumVector := func(a VectorE4) bool { var sum E4 computed := a.Sum() for i := range len(a) { @@ -345,7 +345,7 @@ func TestVectorOps(t *testing.T) { return sum.Equal(&computed) } - innerProductVector := func(a, b Vector) bool { + innerProductVector := func(a, b VectorE4) bool { computed := a.InnerProduct(b) var innerProduct E4 for i := range len(a) { @@ -357,8 +357,8 @@ func TestVectorOps(t *testing.T) { return innerProduct.Equal(&computed) } - mulVector := func(a, b Vector) bool { - c := make(Vector, len(a)) + mulVector := func(a, b VectorE4) bool { + c := make(VectorE4, len(a)) a[0].B0.A0.SetUint64(0x24) b[0].B0.A0.SetUint64(0x42) c.Mul(a, b) @@ -423,8 +423,8 @@ func TestVectorOps(t *testing.T) { )) properties.Property(fmt.Sprintf("vector scalar multiplication by element %d - %s", size, gp.label), prop.ForAll( - func(a Vector, b fr.Element) bool { - c := make(Vector, len(a)) + func(a VectorE4, b fr.Element) bool { + c := make(VectorE4, len(a)) c.ScalarMulByElement(a, &b) for i := range len(a) { var tmp E4 @@ -440,8 +440,8 @@ func TestVectorOps(t *testing.T) { )) properties.Property(fmt.Sprintf("vector multiplication by element %d - %s", size, gp.label), prop.ForAll( - func(a Vector, b fr.Vector) bool { - c := make(Vector, len(a)) + func(a VectorE4, b fr.Vector) bool { + c := make(VectorE4, len(a)) c.MulByElement(a, b) for i := range len(a) { var tmp E4 @@ -458,12 +458,12 @@ func TestVectorOps(t *testing.T) { // checking that in-place butterfly works as intended; properties.Property(fmt.Sprintf("vector butterfly %d - %s", size, gp.label), prop.ForAll( - func(a, b Vector) bool { + func(a, b VectorE4) bool { if len(a) != len(b) { return false } - c := make(Vector, len(a)) - d := make(Vector, len(a)) + c := make(VectorE4, len(a)) + d := make(VectorE4, len(a)) copy(c, a) copy(d, b) c.Butterfly(d) @@ -485,11 +485,11 @@ func TestVectorOps(t *testing.T) { )) properties.Property(fmt.Sprintf("vector butterfly pair %d - %s", size, gp.label), prop.ForAll( - func(a Vector) bool { + func(a VectorE4) bool { if len(a)%2 != 0 { return true // skip odd-sized vectors } - c := make(Vector, len(a)) + c := make(VectorE4, len(a)) copy(c, a) c.ButterflyPair() for i := 0; i < len(a); i += 2 { @@ -507,7 +507,7 @@ func TestVectorOps(t *testing.T) { )) properties.Property(fmt.Sprintf("vector inner product by element %d - %s", size, gp.label), prop.ForAll( - func(a Vector, b fr.Vector) bool { + func(a VectorE4, b fr.Vector) bool { computed := a.InnerProductByElement(b) var innerProduct E4 for i := range len(a) { @@ -527,23 +527,23 @@ func TestVectorOps(t *testing.T) { properties.TestingRun(t, gopter.NewFormatedReporter(false, 260, os.Stdout)) } -// TestVectorExp tests the Exp method for Vector type. +// TestVectorExp tests the Exp method for VectorE4 type. func TestVectorExp(t *testing.T) { assert := require.New(t) // Test with empty vector - empty := make(Vector, 0) - expEmpty := make(Vector, 0) + empty := make(VectorE4, 0) + expEmpty := make(VectorE4, 0) expEmpty.Exp(empty, 5) assert.Equal(0, len(expEmpty), "Exp of empty vector should be empty") // Test with vector of ones and exponent 0 const size = 32 - v := make(Vector, size) + v := make(VectorE4, size) for i := range v { v[i].SetOne() } - expZero := make(Vector, size) + expZero := make(VectorE4, size) expZero.Exp(v, 0) for i := range expZero { assert.True(expZero[i].IsOne(), "Exp(x, 0) should be one for all elements") @@ -553,14 +553,14 @@ func TestVectorExp(t *testing.T) { for i := range v { v[i].MustSetRandom() } - expOne := make(Vector, size) + expOne := make(VectorE4, size) expOne.Exp(v, 1) for i := range v { assert.True(expOne[i].Equal(&v[i]), "Exp(x, 1) should be x for all elements") } // Test with random vector and exponent 2 - expTwo := make(Vector, size) + expTwo := make(VectorE4, size) expTwo.Exp(v, 2) for i := range v { var sq E4 @@ -570,7 +570,7 @@ func TestVectorExp(t *testing.T) { // Test with random vector and exponent k k := int64(7) - expK := make(Vector, size) + expK := make(VectorE4, size) expK.Exp(v, k) for i := range v { var mul E4 @@ -582,7 +582,7 @@ func TestVectorExp(t *testing.T) { } // Test to check v.Exp(v, k) is correct (no modification of v during the process) - vCopy := make(Vector, size) + vCopy := make(VectorE4, size) copy(vCopy, v) vCopy.Exp(vCopy, k) for i := range v { @@ -590,7 +590,7 @@ func TestVectorExp(t *testing.T) { } // Test with random vector and negative exponent -1 - expNegOne := make(Vector, size) + expNegOne := make(VectorE4, size) expNegOne.Exp(v, -1) for i := range v { var inv E4 @@ -601,7 +601,7 @@ func TestVectorExp(t *testing.T) { } // prefixProductGeneric computes the prefix product of the vector in place (single-threaded). -func prefixProductGeneric(vector Vector) { +func prefixProductGeneric(vector VectorE4) { if len(vector) == 0 { return } @@ -610,8 +610,8 @@ func prefixProductGeneric(vector Vector) { } } -func randomVector(size int) Vector { - v := make(Vector, size) +func randomVector(size int) VectorE4 { + v := make(VectorE4, size) for i := range v { v[i].MustSetRandom() } @@ -620,8 +620,8 @@ func randomVector(size int) Vector { func TestPrefixProduct_EmptyVector(t *testing.T) { assert := require.New(t) - v := make(Vector, 0) - expected := make(Vector, 0) + v := make(VectorE4, 0) + expected := make(VectorE4, 0) prefixProductGeneric(expected) v.PrefixProduct() assert.Equal(expected, v) @@ -634,7 +634,7 @@ func TestPrefixProduct_VariousNbTasks(t *testing.T) { for _, size := range sizes { for _, nbTasks := range nbTasksList { v := randomVector(size) - expected := make(Vector, size) + expected := make(VectorE4, size) copy(expected, v) prefixProductGeneric(expected) v.PrefixProduct(nbTasks) @@ -648,8 +648,8 @@ func TestVectorEmptyOps(t *testing.T) { var sum, inner, scalar E4 scalar.MustSetRandom() - empty := make(Vector, 0) - result := make(Vector, 0) + empty := make(VectorE4, 0) + result := make(VectorE4, 0) assert.NotPanics(func() { result.Add(empty, empty) }) assert.NotPanics(func() { result.Sub(empty, empty) }) @@ -665,12 +665,12 @@ func TestVectorEmptyOps(t *testing.T) { func TestVectorSort(t *testing.T) { assert := require.New(t) - v := make(Vector, 3) + v := make(VectorE4, 3) v[0].B0.A0.SetUint64(2) v[1].B0.A0.SetUint64(3) v[2].B0.A0.SetUint64(1) - expected := make(Vector, 3) + expected := make(VectorE4, 3) expected[0].B0.A0.SetUint64(1) expected[1].B0.A0.SetUint64(2) expected[2].B0.A0.SetUint64(3) @@ -685,7 +685,7 @@ func TestVectorSort(t *testing.T) { func TestVectorRoundTrip(t *testing.T) { assert := require.New(t) - v1 := make(Vector, 3) + v1 := make(VectorE4, 3) v1[0].MustSetRandom() v1[1].MustSetRandom() v1[2].MustSetRandom() @@ -693,7 +693,7 @@ func TestVectorRoundTrip(t *testing.T) { b, err := v1.MarshalBinary() assert.NoError(err) - var v2, v3 Vector + var v2, v3 VectorE4 err = v2.UnmarshalBinary(b) assert.NoError(err) @@ -708,12 +708,12 @@ func TestVectorRoundTrip(t *testing.T) { func TestVectorEmptyRoundTrip(t *testing.T) { assert := require.New(t) - v1 := make(Vector, 0) + v1 := make(VectorE4, 0) b, err := v1.MarshalBinary() assert.NoError(err) - var v2, v3 Vector + var v2, v3 VectorE4 err = v2.UnmarshalBinary(b) assert.NoError(err) @@ -728,7 +728,7 @@ func TestVectorEmptyRoundTrip(t *testing.T) { func TestVectorReuseSliceDeserialization(t *testing.T) { assert := require.New(t) - v1 := make(Vector, 4) + v1 := make(VectorE4, 4) for i := range v1 { v1[i].MustSetRandom() } @@ -736,7 +736,7 @@ func TestVectorReuseSliceDeserialization(t *testing.T) { assert.NoError(err) const capacity = 16 - v2 := make(Vector, capacity) + v2 := make(VectorE4, capacity) n, err, errCh := v2.AsyncReadFrom(bytes.NewReader(buf)) assert.Equal(int64(len(buf)), n) assert.NoError(err) @@ -745,7 +745,7 @@ func TestVectorReuseSliceDeserialization(t *testing.T) { assert.Equal(capacity, cap(v2)) assert.True(reflect.DeepEqual(v1, v2)) - v3 := make(Vector, capacity) + v3 := make(VectorE4, capacity) n, err = v3.ReadFrom(bytes.NewReader(buf)) assert.Equal(int64(len(buf)), n) assert.NoError(err) @@ -760,12 +760,12 @@ func TestVectorReadTamperedHeader(t *testing.T) { var input [4]byte binary.BigEndian.PutUint32(input[:], 1<<12) - newVector := func() Vector { - v := make(Vector, 1) + newVector := func() VectorE4 { + v := make(VectorE4, 1) v[0].SetOne() return v } - assertUnchanged := func(v Vector) { + assertUnchanged := func(v VectorE4) { assert.Len(v, 1) assert.Equal(1, cap(v)) var one E4 @@ -796,7 +796,7 @@ func TestVectorReadTamperedHeader(t *testing.T) { func TestVectorReadTamperedHeaderWithoutLen(t *testing.T) { assert := require.New(t) - v1 := make(Vector, 4) + v1 := make(VectorE4, 4) for i := range v1 { v1[i].MustSetRandom() } @@ -811,12 +811,12 @@ func TestVectorReadTamperedHeaderWithoutLen(t *testing.T) { r := readerWithoutLen() _, hasLen := r.(interface{ Len() int }) assert.False(hasLen) - var v2 Vector + var v2 VectorE4 n, err := v2.ReadFrom(r) assert.Equal(int64(len(buf)), n) assert.Error(err) - var v3 Vector + var v3 VectorE4 n, err, errCh := v3.AsyncReadFrom(readerWithoutLen()) assert.Equal(int64(len(buf)), n) assert.Error(err) @@ -824,7 +824,7 @@ func TestVectorReadTamperedHeaderWithoutLen(t *testing.T) { assert.False(open) } -func (vector *Vector) unmarshalBinaryAsync(data []byte) error { +func (vector *VectorE4) unmarshalBinaryAsync(data []byte) error { r := bytes.NewReader(data) _, err, chErr := vector.AsyncReadFrom(r) if err != nil { @@ -918,9 +918,9 @@ func BenchmarkVectorOps(b *testing.B) { // note; to benchmark against "no asm" version, use the following // build tag: -tags purego const N = 1 << 20 - a1 := make(Vector, N) - b1 := make(Vector, N) - c1 := make(Vector, N) + a1 := make(VectorE4, N) + b1 := make(VectorE4, N) + c1 := make(VectorE4, N) b2 := make(fr.Vector, N) for i := 1; i < N; i++ { a1[i-1].MustSetRandom() @@ -1001,7 +1001,7 @@ func BenchmarkVectorOps(b *testing.B) { func BenchmarkPrefixProduct(b *testing.B) { const N = 1 << 19 - a1 := make(Vector, N) + a1 := make(VectorE4, N) for i := range N { a1[i].MustSetRandom() } @@ -1024,7 +1024,7 @@ func BenchmarkPrefixProduct(b *testing.B) { func BenchmarkVectorSerialization(b *testing.B) { const N = 1 << 15 - a1 := make(Vector, N) + a1 := make(VectorE4, N) for i := 1; i < N; i++ { a1[i-1].MustSetRandom() } @@ -1044,7 +1044,7 @@ func BenchmarkVectorSerialization(b *testing.B) { } b.Run("UnmarshalBinary", func(b *testing.B) { - var a2 Vector + var a2 VectorE4 b.ResetTimer() for range b.N { err := a2.UnmarshalBinary(data) @@ -1055,7 +1055,7 @@ func BenchmarkVectorSerialization(b *testing.B) { }) b.Run("unmarshalBinaryAsync", func(b *testing.B) { - var a2 Vector + var a2 VectorE4 b.ResetTimer() for range b.N { err := a2.unmarshalBinaryAsync(data) @@ -1068,7 +1068,7 @@ func BenchmarkVectorSerialization(b *testing.B) { func genZeroVector(size int) gopter.Gen { return func(*gopter.GenParameters) *gopter.GenResult { - return gopter.NewGenResult(make(Vector, size), gopter.NoShrinker) + return gopter.NewGenResult(make(VectorE4, size), gopter.NoShrinker) } } @@ -1076,7 +1076,7 @@ func genMaxVector(size int) gopter.Gen { return func(*gopter.GenParameters) *gopter.GenResult { qMinusOne := fr.Element{2013265921} qMinusOne[0]-- - v := make(Vector, size) + v := make(VectorE4, size) for i := range v { v[i].B0.A0 = qMinusOne v[i].B0.A1 = qMinusOne @@ -1089,7 +1089,7 @@ func genMaxVector(size int) gopter.Gen { func genVector(size int) gopter.Gen { return func(genParams *gopter.GenParameters) *gopter.GenResult { - v := make(Vector, size) + v := make(VectorE4, size) gen := genE4() for i := range v { val, ok := gen(genParams).Retrieve() @@ -1159,3 +1159,35 @@ func TestE4MarshalSetBytesRoundTrip(t *testing.T) { } } } + +func TestE4SetInt64SetUint64(t *testing.T) { + for _, v := range []int64{0, 1, -1, 7, -12345, 1 << 40, -(1 << 40)} { + var z E4 + z.SetInt64(v) + require.Truef(t, z.B0.A1.IsZero(), "coordinate B0.A1 is non-zero for %d", v) + require.Truef(t, z.B1.A0.IsZero(), "coordinate B1.A0 is non-zero for %d", v) + require.Truef(t, z.B1.A1.IsZero(), "coordinate B1.A1 is non-zero for %d", v) + var want fr.Element + want.SetInt64(v) + require.Truef(t, z.B0.A0.Equal(&want), "SetInt64(%d)", v) + + // SetInt64 and SetBigInt agree + var zb E4 + zb.SetBigInt(big.NewInt(v)) + require.Truef(t, z.Equal(&zb), "SetInt64(%d) != SetBigInt(%d)", v, v) + } + for _, v := range []uint64{0, 1, 7, 12345, 1 << 40, 1<<64 - 1} { + var z E4 + z.SetUint64(v) + require.Truef(t, z.B0.A1.IsZero(), "coordinate B0.A1 is non-zero for %d", v) + require.Truef(t, z.B1.A0.IsZero(), "coordinate B1.A0 is non-zero for %d", v) + require.Truef(t, z.B1.A1.IsZero(), "coordinate B1.A1 is non-zero for %d", v) + var want fr.Element + want.SetUint64(v) + require.Truef(t, z.B0.A0.Equal(&want), "SetUint64(%d)", v) + + var zb E4 + zb.SetBigInt(new(big.Int).SetUint64(v)) + require.Truef(t, z.Equal(&zb), "SetUint64(%d) != SetBigInt(%d)", v, v) + } +} diff --git a/field/babybear/extensions/vector.go b/field/babybear/extensions/vector.go index 4991e9d2ab..bfdaaaee89 100644 --- a/field/babybear/extensions/vector.go +++ b/field/babybear/extensions/vector.go @@ -23,10 +23,15 @@ import ( fr "github.com/consensys/gnark-crypto/field/babybear" ) -// Vector represents a vector of E4 elements -type Vector []E4 +// VectorE4 represents a vector of E4 elements +type VectorE4 []E4 -func (vector Vector) Add(a, b Vector) { +// Vector is a slice of E4 elements. +// +// Deprecated: use VectorE4. +type Vector = VectorE4 + +func (vector VectorE4) Add(a, b VectorE4) { N := len(a) if N != len(b) || N != len(vector) { panic("vector.Add: vectors don't have the same length") @@ -34,7 +39,7 @@ func (vector Vector) Add(a, b Vector) { vectorAddGeneric(vector, a, b) } -func (vector Vector) Sub(a, b Vector) { +func (vector VectorE4) Sub(a, b VectorE4) { N := len(a) if N != len(b) || N != len(vector) { panic("vector.Sub: vectors don't have the same length") @@ -42,7 +47,7 @@ func (vector Vector) Sub(a, b Vector) { vectorSubGeneric(vector, a, b) } -func (vector Vector) Mul(a, b Vector) { +func (vector VectorE4) Mul(a, b VectorE4) { N := len(a) if N != len(b) || N != len(vector) { panic("vector.Mul: vectors don't have the same length") @@ -50,7 +55,7 @@ func (vector Vector) Mul(a, b Vector) { vectorMulGeneric(vector, a, b) } -func (vector Vector) ScalarMul(a Vector, b *E4) { +func (vector VectorE4) ScalarMul(a VectorE4, b *E4) { N := len(a) if N != len(vector) { panic("vector.ScalarMul: vectors don't have the same length") @@ -59,11 +64,11 @@ func (vector Vector) ScalarMul(a Vector, b *E4) { } // Sum computes the sum of all elements in the vector. -func (vector Vector) Sum() E4 { +func (vector VectorE4) Sum() E4 { return vectorSumGeneric(vector) } -func (vector Vector) InnerProductByElement(a fr.Vector) E4 { +func (vector VectorE4) InnerProductByElement(a fr.Vector) E4 { N := len(vector) if len(a) != N { panic("vector.InnerProduct: vectors don't have the same length") @@ -71,7 +76,7 @@ func (vector Vector) InnerProductByElement(a fr.Vector) E4 { return vectorInnerProductByElementGeneric(vector, a) } -func (vector Vector) InnerProduct(a Vector) E4 { +func (vector VectorE4) InnerProduct(a VectorE4) E4 { N := len(vector) if len(a) != N { panic("vector.InnerProduct: vectors don't have the same length") @@ -79,7 +84,7 @@ func (vector Vector) InnerProduct(a Vector) E4 { return vectorInnerProductGeneric(vector, a) } -func (vector Vector) MulByElement(a Vector, b fr.Vector) { +func (vector VectorE4) MulByElement(a VectorE4, b fr.Vector) { N := len(vector) if len(a) != N || len(b) != N { panic("vector.MulByElement: vectors don't have the same length") @@ -89,7 +94,7 @@ func (vector Vector) MulByElement(a Vector, b fr.Vector) { // Butterfly computes the in-place butterfly operation on two vectors of E4 elements // If other overlaps with vector, result is undefined, caller should use a temp vector. -func (vector Vector) Butterfly(other Vector) { +func (vector VectorE4) Butterfly(other VectorE4) { N := len(other) if N != len(vector) { panic("vector.Butterfly: vectors don't have the same length") @@ -99,7 +104,7 @@ func (vector Vector) Butterfly(other Vector) { // ButterflyPair computes the in-place butterfly operation of each pair in the vector // vector[0], vector[1]; vector[2], vector[3]; ... -func (vector Vector) ButterflyPair() { +func (vector VectorE4) ButterflyPair() { N := len(vector) if N%2 != 0 { panic("vector.ButterflyPair: vector length must be even") @@ -109,7 +114,7 @@ func (vector Vector) ButterflyPair() { } } -func (vector Vector) ScalarMulByElement(a Vector, b *fr.Element) { +func (vector VectorE4) ScalarMulByElement(a VectorE4, b *fr.Element) { if len(a) != len(vector) { panic("vector.ScalarMulByElement: vectors don't have the same length") } @@ -126,7 +131,7 @@ func (vector Vector) ScalarMulByElement(a Vector, b *fr.Element) { } // Exp sets vector[i] = a[i]ᵏ for all i -func (vector Vector) Exp(a Vector, k int64) { +func (vector VectorE4) Exp(a VectorE4, k int64) { N := len(a) if N != len(vector) { panic("vector.Exp: vectors don't have the same length") @@ -148,7 +153,7 @@ func (vector Vector) Exp(a Vector, k int64) { v0 := &vector[0] // #nosec G602 we check that N > 0 above a0 := &a[0] // #nosec G602 we check that N > 0 above if v0 == a0 { - base = make(Vector, N) + base = make(VectorE4, N) copy(base, a) } } @@ -166,7 +171,7 @@ func (vector Vector) Exp(a Vector, k int64) { // MulAccByElement multiplies each element of the vector v by the E4 element alpha, // accumulating the result in the same vector. -func (vector Vector) MulAccByElement(scale []fr.Element, alpha *E4) { +func (vector VectorE4) MulAccByElement(scale []fr.Element, alpha *E4) { N := len(vector) if N != len(scale) { panic("MulAccByElement: len(vector) != len(scale)") @@ -175,28 +180,28 @@ func (vector Vector) MulAccByElement(scale []fr.Element, alpha *E4) { } // Equal checks whether two vectors are equal -func (vector Vector) Equal(other Vector) bool { +func (vector VectorE4) Equal(other VectorE4) bool { return slices.Equal(vector, other) } // Len is the number of elements in the collection. -func (vector Vector) Len() int { +func (vector VectorE4) Len() int { return len(vector) } // Less reports whether the element with // index i should sort before the element with index j. -func (vector Vector) Less(i, j int) bool { +func (vector VectorE4) Less(i, j int) bool { return vector[i].Cmp(&vector[j]) == -1 } // Swap swaps the elements with indexes i and j. -func (vector Vector) Swap(i, j int) { +func (vector VectorE4) Swap(i, j int) { vector[i], vector[j] = vector[j], vector[i] } // String implements fmt.Stringer interface -func (vector Vector) String() string { +func (vector VectorE4) String() string { var sbb strings.Builder sbb.Grow(len(vector) * 16) sbb.WriteByte('[') @@ -211,7 +216,7 @@ func (vector Vector) String() string { } // MarshalBinary implements encoding.BinaryMarshaler -func (vector *Vector) MarshalBinary() (data []byte, err error) { +func (vector *VectorE4) MarshalBinary() (data []byte, err error) { var buf bytes.Buffer if _, err = vector.WriteTo(&buf); err != nil { @@ -221,7 +226,7 @@ func (vector *Vector) MarshalBinary() (data []byte, err error) { } // UnmarshalBinary implements encoding.BinaryUnmarshaler -func (vector *Vector) UnmarshalBinary(data []byte) error { +func (vector *VectorE4) UnmarshalBinary(data []byte) error { r := bytes.NewReader(data) _, err := vector.ReadFrom(r) return err @@ -229,7 +234,7 @@ func (vector *Vector) UnmarshalBinary(data []byte) error { // WriteTo implements io.WriterTo and writes a vector of big endian encoded Element. // Length of the vector is encoded as a uint32 on the first 4 bytes. -func (vector *Vector) WriteTo(w io.Writer) (int64, error) { +func (vector *VectorE4) WriteTo(w io.Writer) (int64, error) { // encode slice length if err := binary.Write(w, binary.BigEndian, uint32(len(*vector))); err != nil { @@ -257,7 +262,7 @@ func (vector *Vector) WriteTo(w io.Writer) (int64, error) { return n, nil } -// AsyncReadFrom implements an asynchronous version of [Vector.ReadFrom]. It +// AsyncReadFrom implements an asynchronous version of [VectorE4.ReadFrom]. It // reads the reader r in full and then performs the validation and conversion to // Montgomery form separately in a goroutine. Any error encountered during // reading is returned directly, while errors encountered during @@ -281,7 +286,7 @@ func (vector *Vector) WriteTo(w io.Writer) (int64, error) { // - first 4 bytes: length of the vector as a big-endian uint32 // - for each element of the vector, `4 * fr.Bytes` bytes representing the // element in big-endian encoding. -func (vector *Vector) AsyncReadFrom(r io.Reader) (int64, error, chan error) { // nolint ST1008 +func (vector *VectorE4) AsyncReadFrom(r io.Reader) (int64, error, chan error) { // nolint ST1008 chErr := make(chan error, 1) var bufSizeSlice [4]byte if read, err := io.ReadFull(r, bufSizeSlice[:]); err != nil { @@ -310,7 +315,7 @@ func (vector *Vector) AsyncReadFrom(r io.Reader) (int64, error, chan error) { // *vector = (*vector)[:0] if headerSliceLen == 0 { if *vector == nil { - *vector = Vector{} + *vector = VectorE4{} } close(chErr) return totalRead, nil, chErr @@ -318,7 +323,7 @@ func (vector *Vector) AsyncReadFrom(r io.Reader) (int64, error, chan error) { // for i := uint64(0); i < headerSliceLen; i += maxAllocateSliceLength { if len(*vector) <= int(i) { - *vector = append(*vector, make(Vector, int(min(headerSliceLen-i, maxAllocateSliceLength)))...) + *vector = append(*vector, make(VectorE4, int(min(headerSliceLen-i, maxAllocateSliceLength)))...) } bSlice := unsafe.Slice((*byte)(unsafe.Pointer(&(*vector)[i])), int(min(headerSliceLen-i, maxAllocateSliceLength))*BytesE4) read, err := io.ReadFull(r, bSlice) @@ -391,7 +396,7 @@ func (vector *Vector) AsyncReadFrom(r io.Reader) (int64, error, chan error) { // // - for each element of the vector, `4 * fr.Bytes` bytes representing the element in big-endian encoding. // // The method implements [io.ReaderFrom] interface. -func (vector *Vector) ReadFrom(r io.Reader) (int64, error) { +func (vector *VectorE4) ReadFrom(r io.Reader) (int64, error) { // call the async version and wait for the channel to be closed n, err, chErr := vector.AsyncReadFrom(r) @@ -405,7 +410,7 @@ func (vector *Vector) ReadFrom(r io.Reader) (int64, error) { // i.e. vector[i] = vector[0] * vector[1] * ... * vector[i] // If nbTasks > 1, it uses nbTasks goroutines to compute the prefix product in parallel. // If nbTasks is not provided, it uses the number of CPU cores. -func (vector Vector) PrefixProduct(nbTasks ...int) { +func (vector VectorE4) PrefixProduct(nbTasks ...int) { N := len(vector) if N < 2 { return @@ -465,34 +470,34 @@ func (vector Vector) PrefixProduct(nbTasks ...int) { } -func (vector Vector) prefixProductGeneric() { +func (vector VectorE4) prefixProductGeneric() { for i := 1; i < len(vector); i++ { vector[i].Mul(&vector[i], &vector[i-1]) } } -func vectorAddGeneric(res, a, b Vector) { +func vectorAddGeneric(res, a, b VectorE4) { for i := range len(res) { res[i].Add(&a[i], &b[i]) } } -func vectorSubGeneric(res, a, b Vector) { +func vectorSubGeneric(res, a, b VectorE4) { for i := range len(res) { res[i].Sub(&a[i], &b[i]) } } -func vectorMulGeneric(res, a, b Vector) { +func vectorMulGeneric(res, a, b VectorE4) { for i := range len(res) { res[i].Mul(&a[i], &b[i]) } } -func vectorScalarMulGeneric(res, a Vector, b *E4) { +func vectorScalarMulGeneric(res, a VectorE4, b *E4) { for i := range len(res) { res[i].Mul(&a[i], b) } } -func vectorInnerProductGeneric(a, b Vector) E4 { +func vectorInnerProductGeneric(a, b VectorE4) E4 { var res, tmp E4 for i := range len(a) { tmp.Mul(&a[i], &b[i]) @@ -501,7 +506,7 @@ func vectorInnerProductGeneric(a, b Vector) E4 { return res } -func vectorInnerProductByElementGeneric(a Vector, b fr.Vector) E4 { +func vectorInnerProductByElementGeneric(a VectorE4, b fr.Vector) E4 { var res, tmp E4 for i := range len(a) { tmp.MulByElement(&a[i], &b[i]) @@ -510,7 +515,7 @@ func vectorInnerProductByElementGeneric(a Vector, b fr.Vector) E4 { return res } -func vectorSumGeneric(v Vector) E4 { +func vectorSumGeneric(v VectorE4) E4 { var sum E4 for i := range len(v) { sum.Add(&sum, &v[i]) @@ -518,7 +523,7 @@ func vectorSumGeneric(v Vector) E4 { return sum } -func vectorMulAccByElementGeneric(v Vector, scale []fr.Element, alpha *E4) { +func vectorMulAccByElementGeneric(v VectorE4, scale []fr.Element, alpha *E4) { var tmp E4 for i := range len(v) { tmp.MulByElement(alpha, &scale[i]) @@ -526,13 +531,13 @@ func vectorMulAccByElementGeneric(v Vector, scale []fr.Element, alpha *E4) { } } -func vectorMulByElementGeneric(res, a Vector, b fr.Vector) { +func vectorMulByElementGeneric(res, a VectorE4, b fr.Vector) { for i := range len(res) { res[i].MulByElement(&a[i], &b[i]) } } -func vectorButterflyGeneric(a, b Vector) { +func vectorButterflyGeneric(a, b VectorE4) { for i := range len(a) { Butterfly(&a[i], &b[i]) } diff --git a/field/babybear/fft/fftext.go b/field/babybear/fft/fftext.go index 1e6ace9a83..87ad123eba 100644 --- a/field/babybear/fft/fftext.go +++ b/field/babybear/fft/fftext.go @@ -57,7 +57,7 @@ func (domain *Domain) FFTExt(a []fext.E4, decimation Decimation, opts ...Option) } } parallel.ExecuteAligned(len(a), 4, func(start, end int) { - va := fext.Vector(a[start:end]) + va := fext.VectorE4(a[start:end]) va.MulByElement(va, cosetTable[start:end]) }, opt.nbTasks) } @@ -126,7 +126,7 @@ func (domain *Domain) FFTInverseExt(a []fext.E4, decimation Decimation, opts ... // scale by CardinalityInv if !opt.coset { parallel.ExecuteAligned(len(a), 4, func(start, end int) { - va := fext.Vector(a[start:end]) + va := fext.VectorE4(a[start:end]) va.ScalarMulByElement(va, &domain.CardinalityInv) }, opt.nbTasks) return @@ -156,7 +156,7 @@ func (domain *Domain) FFTInverseExt(a []fext.E4, decimation Decimation, opts ... } } parallel.ExecuteAligned(len(a), 4, func(start, end int) { - va := fext.Vector(a[start:end]) + va := fext.VectorE4(a[start:end]) va.MulByElement(va, cosetTableInv[start:end]) va.ScalarMulByElement(va, &domain.CardinalityInv) }, opt.nbTasks) @@ -226,8 +226,8 @@ func difFFTExt(a []fext.E4, w babybear.Element, twiddles [][]babybear.Element, t } func innerDIFWithTwiddlesExt(a []fext.E4, twiddles []babybear.Element, start, end, m int) { - va0 := fext.Vector(a[start:end]) - va1 := fext.Vector(a[start+m : end+m]) + va0 := fext.VectorE4(a[start:end]) + va1 := fext.VectorE4(a[start+m : end+m]) va0.Butterfly(va1) va1.MulByElement(va1, twiddles[start:end]) } @@ -308,8 +308,8 @@ func ditFFTExt(a []fext.E4, w babybear.Element, twiddles [][]babybear.Element, t } func innerDITWithTwiddlesExt(a []fext.E4, twiddles []babybear.Element, start, end, m int) { - va0 := fext.Vector(a[start:end]) - va1 := fext.Vector(a[start+m : end+m]) + va0 := fext.VectorE4(a[start:end]) + va1 := fext.VectorE4(a[start+m : end+m]) va1.MulByElement(va1, twiddles[start:end]) va0.Butterfly(va1) } @@ -357,14 +357,14 @@ func kerDIFNP_512Ext(a []fext.E4, twiddles [][]babybear.Element, stage int) { for offset := 0; offset < 512; offset += 4 { innerDIFWithTwiddlesExt(a[offset:offset+4], twiddles[stage+7], 0, 2, 2) } - va := fext.Vector(a[:512]) + va := fext.VectorE4(a[:512]) va.ButterflyPair() } func kerDITNP_512Ext(a []fext.E4, twiddles [][]babybear.Element, stage int) { // code unrolled & generated by internal/generator/fft/template/fftext.go.tmpl - va := fext.Vector(a[:512]) + va := fext.VectorE4(a[:512]) va.ButterflyPair() for offset := 0; offset < 512; offset += 4 { innerDITWithTwiddlesExtM2(a[offset:offset+4], twiddles[stage+7]) diff --git a/field/goldilocks/extensions/e2.go b/field/goldilocks/extensions/e2.go index a4079c0634..5ea33e19fa 100644 --- a/field/goldilocks/extensions/e2.go +++ b/field/goldilocks/extensions/e2.go @@ -75,6 +75,20 @@ func (z *E2) SetOne() *E2 { return z } +// SetInt64 sets z to v (embedded via the unique ring homomorphism ℤ→𝔽) and returns z +func (z *E2) SetInt64(v int64) *E2 { + *z = E2{} + z.A0.SetInt64(v) + return z +} + +// SetUint64 sets z to v (embedded via the unique ring homomorphism ℤ→𝔽) and returns z +func (z *E2) SetUint64(v uint64) *E2 { + *z = E2{} + z.A0.SetUint64(v) + return z +} + // SetBigInt sets z to v (embedded via the unique ring homomorphism ℤ→𝔽, reduced modulo the base field order) and returns z func (z *E2) SetBigInt(v *big.Int) *E2 { *z = E2{} diff --git a/field/goldilocks/extensions/e2_test.go b/field/goldilocks/extensions/e2_test.go index e14bf22aa7..7c10621772 100644 --- a/field/goldilocks/extensions/e2_test.go +++ b/field/goldilocks/extensions/e2_test.go @@ -572,3 +572,31 @@ func TestE2MarshalSetBytesRoundTrip(t *testing.T) { } } } + +func TestE2SetInt64SetUint64(t *testing.T) { + for _, v := range []int64{0, 1, -1, 7, -12345, 1 << 40, -(1 << 40)} { + var z E2 + z.SetInt64(v) + require.Truef(t, z.A1.IsZero(), "coordinate A1 is non-zero for %d", v) + var want fr.Element + want.SetInt64(v) + require.Truef(t, z.A0.Equal(&want), "SetInt64(%d)", v) + + // SetInt64 and SetBigInt agree + var zb E2 + zb.SetBigInt(big.NewInt(v)) + require.Truef(t, z.Equal(&zb), "SetInt64(%d) != SetBigInt(%d)", v, v) + } + for _, v := range []uint64{0, 1, 7, 12345, 1 << 40, 1<<64 - 1} { + var z E2 + z.SetUint64(v) + require.Truef(t, z.A1.IsZero(), "coordinate A1 is non-zero for %d", v) + var want fr.Element + want.SetUint64(v) + require.Truef(t, z.A0.Equal(&want), "SetUint64(%d)", v) + + var zb E2 + zb.SetBigInt(new(big.Int).SetUint64(v)) + require.Truef(t, z.Equal(&zb), "SetUint64(%d) != SetBigInt(%d)", v, v) + } +} diff --git a/field/koalabear/extensions/e2.go b/field/koalabear/extensions/e2.go index b6d18b2cb8..fa89720ee7 100644 --- a/field/koalabear/extensions/e2.go +++ b/field/koalabear/extensions/e2.go @@ -75,6 +75,20 @@ func (z *E2) SetOne() *E2 { return z } +// SetInt64 sets z to v (embedded via the unique ring homomorphism ℤ→𝔽) and returns z +func (z *E2) SetInt64(v int64) *E2 { + *z = E2{} + z.A0.SetInt64(v) + return z +} + +// SetUint64 sets z to v (embedded via the unique ring homomorphism ℤ→𝔽) and returns z +func (z *E2) SetUint64(v uint64) *E2 { + *z = E2{} + z.A0.SetUint64(v) + return z +} + // SetBigInt sets z to v (embedded via the unique ring homomorphism ℤ→𝔽, reduced modulo the base field order) and returns z func (z *E2) SetBigInt(v *big.Int) *E2 { *z = E2{} diff --git a/field/koalabear/extensions/e2_test.go b/field/koalabear/extensions/e2_test.go index c2429fb21c..8ee2dd1f18 100644 --- a/field/koalabear/extensions/e2_test.go +++ b/field/koalabear/extensions/e2_test.go @@ -589,3 +589,31 @@ func TestE2MarshalSetBytesRoundTrip(t *testing.T) { } } } + +func TestE2SetInt64SetUint64(t *testing.T) { + for _, v := range []int64{0, 1, -1, 7, -12345, 1 << 40, -(1 << 40)} { + var z E2 + z.SetInt64(v) + require.Truef(t, z.A1.IsZero(), "coordinate A1 is non-zero for %d", v) + var want fr.Element + want.SetInt64(v) + require.Truef(t, z.A0.Equal(&want), "SetInt64(%d)", v) + + // SetInt64 and SetBigInt agree + var zb E2 + zb.SetBigInt(big.NewInt(v)) + require.Truef(t, z.Equal(&zb), "SetInt64(%d) != SetBigInt(%d)", v, v) + } + for _, v := range []uint64{0, 1, 7, 12345, 1 << 40, 1<<64 - 1} { + var z E2 + z.SetUint64(v) + require.Truef(t, z.A1.IsZero(), "coordinate A1 is non-zero for %d", v) + var want fr.Element + want.SetUint64(v) + require.Truef(t, z.A0.Equal(&want), "SetUint64(%d)", v) + + var zb E2 + zb.SetBigInt(new(big.Int).SetUint64(v)) + require.Truef(t, z.Equal(&zb), "SetUint64(%d) != SetBigInt(%d)", v, v) + } +} diff --git a/field/koalabear/extensions/e4.go b/field/koalabear/extensions/e4.go index 9a2a4b2436..c164108368 100644 --- a/field/koalabear/extensions/e4.go +++ b/field/koalabear/extensions/e4.go @@ -86,6 +86,20 @@ func (z *E4) SetOne() *E4 { return z } +// SetInt64 sets z to v (embedded via the unique ring homomorphism ℤ→𝔽) and returns z +func (z *E4) SetInt64(v int64) *E4 { + *z = E4{} + z.B0.A0.SetInt64(v) + return z +} + +// SetUint64 sets z to v (embedded via the unique ring homomorphism ℤ→𝔽) and returns z +func (z *E4) SetUint64(v uint64) *E4 { + *z = E4{} + z.B0.A0.SetUint64(v) + return z +} + // SetBigInt sets z to v (embedded via the unique ring homomorphism ℤ→𝔽, reduced modulo the base field order) and returns z func (z *E4) SetBigInt(v *big.Int) *E4 { *z = E4{} diff --git a/field/koalabear/extensions/e4_test.go b/field/koalabear/extensions/e4_test.go index 46f4bf75a1..d15b36792d 100644 --- a/field/koalabear/extensions/e4_test.go +++ b/field/koalabear/extensions/e4_test.go @@ -293,8 +293,8 @@ func TestVectorOps(t *testing.T) { } properties := gopter.NewProperties(parameters) - addVector := func(a, b Vector) bool { - c := make(Vector, len(a)) + addVector := func(a, b VectorE4) bool { + c := make(VectorE4, len(a)) c.Add(a, b) for i := range len(a) { @@ -307,8 +307,8 @@ func TestVectorOps(t *testing.T) { return true } - subVector := func(a, b Vector) bool { - c := make(Vector, len(a)) + subVector := func(a, b VectorE4) bool { + c := make(VectorE4, len(a)) c.Sub(a, b) for i := range len(a) { @@ -321,8 +321,8 @@ func TestVectorOps(t *testing.T) { return true } - scalarMulVector := func(a Vector, b E4) bool { - c := make(Vector, len(a)) + scalarMulVector := func(a VectorE4, b E4) bool { + c := make(VectorE4, len(a)) c.ScalarMul(a, &b) for i := range len(a) { @@ -335,7 +335,7 @@ func TestVectorOps(t *testing.T) { return true } - sumVector := func(a Vector) bool { + sumVector := func(a VectorE4) bool { var sum E4 computed := a.Sum() for i := range len(a) { @@ -345,7 +345,7 @@ func TestVectorOps(t *testing.T) { return sum.Equal(&computed) } - innerProductVector := func(a, b Vector) bool { + innerProductVector := func(a, b VectorE4) bool { computed := a.InnerProduct(b) var innerProduct E4 for i := range len(a) { @@ -357,8 +357,8 @@ func TestVectorOps(t *testing.T) { return innerProduct.Equal(&computed) } - mulVector := func(a, b Vector) bool { - c := make(Vector, len(a)) + mulVector := func(a, b VectorE4) bool { + c := make(VectorE4, len(a)) a[0].B0.A0.SetUint64(0x24) b[0].B0.A0.SetUint64(0x42) c.Mul(a, b) @@ -423,8 +423,8 @@ func TestVectorOps(t *testing.T) { )) properties.Property(fmt.Sprintf("vector scalar multiplication by element %d - %s", size, gp.label), prop.ForAll( - func(a Vector, b fr.Element) bool { - c := make(Vector, len(a)) + func(a VectorE4, b fr.Element) bool { + c := make(VectorE4, len(a)) c.ScalarMulByElement(a, &b) for i := range len(a) { var tmp E4 @@ -440,8 +440,8 @@ func TestVectorOps(t *testing.T) { )) properties.Property(fmt.Sprintf("vector multiplication by element %d - %s", size, gp.label), prop.ForAll( - func(a Vector, b fr.Vector) bool { - c := make(Vector, len(a)) + func(a VectorE4, b fr.Vector) bool { + c := make(VectorE4, len(a)) c.MulByElement(a, b) for i := range len(a) { var tmp E4 @@ -458,12 +458,12 @@ func TestVectorOps(t *testing.T) { // checking that in-place butterfly works as intended; properties.Property(fmt.Sprintf("vector butterfly %d - %s", size, gp.label), prop.ForAll( - func(a, b Vector) bool { + func(a, b VectorE4) bool { if len(a) != len(b) { return false } - c := make(Vector, len(a)) - d := make(Vector, len(a)) + c := make(VectorE4, len(a)) + d := make(VectorE4, len(a)) copy(c, a) copy(d, b) c.Butterfly(d) @@ -485,11 +485,11 @@ func TestVectorOps(t *testing.T) { )) properties.Property(fmt.Sprintf("vector butterfly pair %d - %s", size, gp.label), prop.ForAll( - func(a Vector) bool { + func(a VectorE4) bool { if len(a)%2 != 0 { return true // skip odd-sized vectors } - c := make(Vector, len(a)) + c := make(VectorE4, len(a)) copy(c, a) c.ButterflyPair() for i := 0; i < len(a); i += 2 { @@ -507,7 +507,7 @@ func TestVectorOps(t *testing.T) { )) properties.Property(fmt.Sprintf("vector inner product by element %d - %s", size, gp.label), prop.ForAll( - func(a Vector, b fr.Vector) bool { + func(a VectorE4, b fr.Vector) bool { computed := a.InnerProductByElement(b) var innerProduct E4 for i := range len(a) { @@ -527,23 +527,23 @@ func TestVectorOps(t *testing.T) { properties.TestingRun(t, gopter.NewFormatedReporter(false, 260, os.Stdout)) } -// TestVectorExp tests the Exp method for Vector type. +// TestVectorExp tests the Exp method for VectorE4 type. func TestVectorExp(t *testing.T) { assert := require.New(t) // Test with empty vector - empty := make(Vector, 0) - expEmpty := make(Vector, 0) + empty := make(VectorE4, 0) + expEmpty := make(VectorE4, 0) expEmpty.Exp(empty, 5) assert.Equal(0, len(expEmpty), "Exp of empty vector should be empty") // Test with vector of ones and exponent 0 const size = 32 - v := make(Vector, size) + v := make(VectorE4, size) for i := range v { v[i].SetOne() } - expZero := make(Vector, size) + expZero := make(VectorE4, size) expZero.Exp(v, 0) for i := range expZero { assert.True(expZero[i].IsOne(), "Exp(x, 0) should be one for all elements") @@ -553,14 +553,14 @@ func TestVectorExp(t *testing.T) { for i := range v { v[i].MustSetRandom() } - expOne := make(Vector, size) + expOne := make(VectorE4, size) expOne.Exp(v, 1) for i := range v { assert.True(expOne[i].Equal(&v[i]), "Exp(x, 1) should be x for all elements") } // Test with random vector and exponent 2 - expTwo := make(Vector, size) + expTwo := make(VectorE4, size) expTwo.Exp(v, 2) for i := range v { var sq E4 @@ -570,7 +570,7 @@ func TestVectorExp(t *testing.T) { // Test with random vector and exponent k k := int64(7) - expK := make(Vector, size) + expK := make(VectorE4, size) expK.Exp(v, k) for i := range v { var mul E4 @@ -582,7 +582,7 @@ func TestVectorExp(t *testing.T) { } // Test to check v.Exp(v, k) is correct (no modification of v during the process) - vCopy := make(Vector, size) + vCopy := make(VectorE4, size) copy(vCopy, v) vCopy.Exp(vCopy, k) for i := range v { @@ -590,7 +590,7 @@ func TestVectorExp(t *testing.T) { } // Test with random vector and negative exponent -1 - expNegOne := make(Vector, size) + expNegOne := make(VectorE4, size) expNegOne.Exp(v, -1) for i := range v { var inv E4 @@ -601,7 +601,7 @@ func TestVectorExp(t *testing.T) { } // prefixProductGeneric computes the prefix product of the vector in place (single-threaded). -func prefixProductGeneric(vector Vector) { +func prefixProductGeneric(vector VectorE4) { if len(vector) == 0 { return } @@ -610,8 +610,8 @@ func prefixProductGeneric(vector Vector) { } } -func randomVector(size int) Vector { - v := make(Vector, size) +func randomVector(size int) VectorE4 { + v := make(VectorE4, size) for i := range v { v[i].MustSetRandom() } @@ -620,8 +620,8 @@ func randomVector(size int) Vector { func TestPrefixProduct_EmptyVector(t *testing.T) { assert := require.New(t) - v := make(Vector, 0) - expected := make(Vector, 0) + v := make(VectorE4, 0) + expected := make(VectorE4, 0) prefixProductGeneric(expected) v.PrefixProduct() assert.Equal(expected, v) @@ -634,7 +634,7 @@ func TestPrefixProduct_VariousNbTasks(t *testing.T) { for _, size := range sizes { for _, nbTasks := range nbTasksList { v := randomVector(size) - expected := make(Vector, size) + expected := make(VectorE4, size) copy(expected, v) prefixProductGeneric(expected) v.PrefixProduct(nbTasks) @@ -648,8 +648,8 @@ func TestVectorEmptyOps(t *testing.T) { var sum, inner, scalar E4 scalar.MustSetRandom() - empty := make(Vector, 0) - result := make(Vector, 0) + empty := make(VectorE4, 0) + result := make(VectorE4, 0) assert.NotPanics(func() { result.Add(empty, empty) }) assert.NotPanics(func() { result.Sub(empty, empty) }) @@ -665,12 +665,12 @@ func TestVectorEmptyOps(t *testing.T) { func TestVectorSort(t *testing.T) { assert := require.New(t) - v := make(Vector, 3) + v := make(VectorE4, 3) v[0].B0.A0.SetUint64(2) v[1].B0.A0.SetUint64(3) v[2].B0.A0.SetUint64(1) - expected := make(Vector, 3) + expected := make(VectorE4, 3) expected[0].B0.A0.SetUint64(1) expected[1].B0.A0.SetUint64(2) expected[2].B0.A0.SetUint64(3) @@ -685,7 +685,7 @@ func TestVectorSort(t *testing.T) { func TestVectorRoundTrip(t *testing.T) { assert := require.New(t) - v1 := make(Vector, 3) + v1 := make(VectorE4, 3) v1[0].MustSetRandom() v1[1].MustSetRandom() v1[2].MustSetRandom() @@ -693,7 +693,7 @@ func TestVectorRoundTrip(t *testing.T) { b, err := v1.MarshalBinary() assert.NoError(err) - var v2, v3 Vector + var v2, v3 VectorE4 err = v2.UnmarshalBinary(b) assert.NoError(err) @@ -708,12 +708,12 @@ func TestVectorRoundTrip(t *testing.T) { func TestVectorEmptyRoundTrip(t *testing.T) { assert := require.New(t) - v1 := make(Vector, 0) + v1 := make(VectorE4, 0) b, err := v1.MarshalBinary() assert.NoError(err) - var v2, v3 Vector + var v2, v3 VectorE4 err = v2.UnmarshalBinary(b) assert.NoError(err) @@ -728,7 +728,7 @@ func TestVectorEmptyRoundTrip(t *testing.T) { func TestVectorReuseSliceDeserialization(t *testing.T) { assert := require.New(t) - v1 := make(Vector, 4) + v1 := make(VectorE4, 4) for i := range v1 { v1[i].MustSetRandom() } @@ -736,7 +736,7 @@ func TestVectorReuseSliceDeserialization(t *testing.T) { assert.NoError(err) const capacity = 16 - v2 := make(Vector, capacity) + v2 := make(VectorE4, capacity) n, err, errCh := v2.AsyncReadFrom(bytes.NewReader(buf)) assert.Equal(int64(len(buf)), n) assert.NoError(err) @@ -745,7 +745,7 @@ func TestVectorReuseSliceDeserialization(t *testing.T) { assert.Equal(capacity, cap(v2)) assert.True(reflect.DeepEqual(v1, v2)) - v3 := make(Vector, capacity) + v3 := make(VectorE4, capacity) n, err = v3.ReadFrom(bytes.NewReader(buf)) assert.Equal(int64(len(buf)), n) assert.NoError(err) @@ -760,12 +760,12 @@ func TestVectorReadTamperedHeader(t *testing.T) { var input [4]byte binary.BigEndian.PutUint32(input[:], 1<<12) - newVector := func() Vector { - v := make(Vector, 1) + newVector := func() VectorE4 { + v := make(VectorE4, 1) v[0].SetOne() return v } - assertUnchanged := func(v Vector) { + assertUnchanged := func(v VectorE4) { assert.Len(v, 1) assert.Equal(1, cap(v)) var one E4 @@ -796,7 +796,7 @@ func TestVectorReadTamperedHeader(t *testing.T) { func TestVectorReadTamperedHeaderWithoutLen(t *testing.T) { assert := require.New(t) - v1 := make(Vector, 4) + v1 := make(VectorE4, 4) for i := range v1 { v1[i].MustSetRandom() } @@ -811,12 +811,12 @@ func TestVectorReadTamperedHeaderWithoutLen(t *testing.T) { r := readerWithoutLen() _, hasLen := r.(interface{ Len() int }) assert.False(hasLen) - var v2 Vector + var v2 VectorE4 n, err := v2.ReadFrom(r) assert.Equal(int64(len(buf)), n) assert.Error(err) - var v3 Vector + var v3 VectorE4 n, err, errCh := v3.AsyncReadFrom(readerWithoutLen()) assert.Equal(int64(len(buf)), n) assert.Error(err) @@ -824,7 +824,7 @@ func TestVectorReadTamperedHeaderWithoutLen(t *testing.T) { assert.False(open) } -func (vector *Vector) unmarshalBinaryAsync(data []byte) error { +func (vector *VectorE4) unmarshalBinaryAsync(data []byte) error { r := bytes.NewReader(data) _, err, chErr := vector.AsyncReadFrom(r) if err != nil { @@ -918,9 +918,9 @@ func BenchmarkVectorOps(b *testing.B) { // note; to benchmark against "no asm" version, use the following // build tag: -tags purego const N = 1 << 20 - a1 := make(Vector, N) - b1 := make(Vector, N) - c1 := make(Vector, N) + a1 := make(VectorE4, N) + b1 := make(VectorE4, N) + c1 := make(VectorE4, N) b2 := make(fr.Vector, N) for i := 1; i < N; i++ { a1[i-1].MustSetRandom() @@ -1001,7 +1001,7 @@ func BenchmarkVectorOps(b *testing.B) { func BenchmarkPrefixProduct(b *testing.B) { const N = 1 << 19 - a1 := make(Vector, N) + a1 := make(VectorE4, N) for i := range N { a1[i].MustSetRandom() } @@ -1024,7 +1024,7 @@ func BenchmarkPrefixProduct(b *testing.B) { func BenchmarkVectorSerialization(b *testing.B) { const N = 1 << 15 - a1 := make(Vector, N) + a1 := make(VectorE4, N) for i := 1; i < N; i++ { a1[i-1].MustSetRandom() } @@ -1044,7 +1044,7 @@ func BenchmarkVectorSerialization(b *testing.B) { } b.Run("UnmarshalBinary", func(b *testing.B) { - var a2 Vector + var a2 VectorE4 b.ResetTimer() for range b.N { err := a2.UnmarshalBinary(data) @@ -1055,7 +1055,7 @@ func BenchmarkVectorSerialization(b *testing.B) { }) b.Run("unmarshalBinaryAsync", func(b *testing.B) { - var a2 Vector + var a2 VectorE4 b.ResetTimer() for range b.N { err := a2.unmarshalBinaryAsync(data) @@ -1068,7 +1068,7 @@ func BenchmarkVectorSerialization(b *testing.B) { func genZeroVector(size int) gopter.Gen { return func(*gopter.GenParameters) *gopter.GenResult { - return gopter.NewGenResult(make(Vector, size), gopter.NoShrinker) + return gopter.NewGenResult(make(VectorE4, size), gopter.NoShrinker) } } @@ -1076,7 +1076,7 @@ func genMaxVector(size int) gopter.Gen { return func(*gopter.GenParameters) *gopter.GenResult { qMinusOne := fr.Element{2130706433} qMinusOne[0]-- - v := make(Vector, size) + v := make(VectorE4, size) for i := range v { v[i].B0.A0 = qMinusOne v[i].B0.A1 = qMinusOne @@ -1089,7 +1089,7 @@ func genMaxVector(size int) gopter.Gen { func genVector(size int) gopter.Gen { return func(genParams *gopter.GenParameters) *gopter.GenResult { - v := make(Vector, size) + v := make(VectorE4, size) gen := genE4() for i := range v { val, ok := gen(genParams).Retrieve() @@ -1159,3 +1159,35 @@ func TestE4MarshalSetBytesRoundTrip(t *testing.T) { } } } + +func TestE4SetInt64SetUint64(t *testing.T) { + for _, v := range []int64{0, 1, -1, 7, -12345, 1 << 40, -(1 << 40)} { + var z E4 + z.SetInt64(v) + require.Truef(t, z.B0.A1.IsZero(), "coordinate B0.A1 is non-zero for %d", v) + require.Truef(t, z.B1.A0.IsZero(), "coordinate B1.A0 is non-zero for %d", v) + require.Truef(t, z.B1.A1.IsZero(), "coordinate B1.A1 is non-zero for %d", v) + var want fr.Element + want.SetInt64(v) + require.Truef(t, z.B0.A0.Equal(&want), "SetInt64(%d)", v) + + // SetInt64 and SetBigInt agree + var zb E4 + zb.SetBigInt(big.NewInt(v)) + require.Truef(t, z.Equal(&zb), "SetInt64(%d) != SetBigInt(%d)", v, v) + } + for _, v := range []uint64{0, 1, 7, 12345, 1 << 40, 1<<64 - 1} { + var z E4 + z.SetUint64(v) + require.Truef(t, z.B0.A1.IsZero(), "coordinate B0.A1 is non-zero for %d", v) + require.Truef(t, z.B1.A0.IsZero(), "coordinate B1.A0 is non-zero for %d", v) + require.Truef(t, z.B1.A1.IsZero(), "coordinate B1.A1 is non-zero for %d", v) + var want fr.Element + want.SetUint64(v) + require.Truef(t, z.B0.A0.Equal(&want), "SetUint64(%d)", v) + + var zb E4 + zb.SetBigInt(new(big.Int).SetUint64(v)) + require.Truef(t, z.Equal(&zb), "SetUint64(%d) != SetBigInt(%d)", v, v) + } +} diff --git a/field/koalabear/extensions/vector.go b/field/koalabear/extensions/vector.go index 9f40b9a8d5..f7541101d2 100644 --- a/field/koalabear/extensions/vector.go +++ b/field/koalabear/extensions/vector.go @@ -24,10 +24,15 @@ import ( "github.com/consensys/gnark-crypto/utils/cpu" ) -// Vector represents a vector of E4 elements -type Vector []E4 +// VectorE4 represents a vector of E4 elements +type VectorE4 []E4 -func (vector Vector) Add(a, b Vector) { +// Vector is a slice of E4 elements. +// +// Deprecated: use VectorE4. +type Vector = VectorE4 + +func (vector VectorE4) Add(a, b VectorE4) { N := len(a) if N != len(b) || N != len(vector) { panic("vector.Add: vectors don't have the same length") @@ -45,7 +50,7 @@ func (vector Vector) Add(a, b Vector) { } } -func (vector Vector) Sub(a, b Vector) { +func (vector VectorE4) Sub(a, b VectorE4) { N := len(a) if N != len(b) || N != len(vector) { panic("vector.Sub: vectors don't have the same length") @@ -63,7 +68,7 @@ func (vector Vector) Sub(a, b Vector) { } } -func (vector Vector) Mul(a, b Vector) { +func (vector VectorE4) Mul(a, b VectorE4) { N := len(a) if N != len(b) || N != len(vector) { panic("vector.Mul: vectors don't have the same length") @@ -81,7 +86,7 @@ func (vector Vector) Mul(a, b Vector) { } } -func (vector Vector) ScalarMul(a Vector, b *E4) { +func (vector VectorE4) ScalarMul(a VectorE4, b *E4) { N := len(a) if N != len(vector) { panic("vector.ScalarMul: vectors don't have the same length") @@ -100,7 +105,7 @@ func (vector Vector) ScalarMul(a Vector, b *E4) { } // Sum computes the sum of all elements in the vector. -func (vector Vector) Sum() E4 { +func (vector VectorE4) Sum() E4 { const blockSize = 2 N := len(vector) if !cpu.SupportAVX512 || N < blockSize { @@ -125,7 +130,7 @@ func (vector Vector) Sum() E4 { return res } -func (vector Vector) InnerProductByElement(a fr.Vector) E4 { +func (vector VectorE4) InnerProductByElement(a fr.Vector) E4 { N := len(vector) if len(a) != N { panic("vector.InnerProduct: vectors don't have the same length") @@ -145,7 +150,7 @@ func (vector Vector) InnerProductByElement(a fr.Vector) E4 { return res } -func (vector Vector) InnerProduct(a Vector) E4 { +func (vector VectorE4) InnerProduct(a VectorE4) E4 { N := len(vector) if len(a) != N { panic("vector.InnerProduct: vectors don't have the same length") @@ -187,7 +192,7 @@ func (vector Vector) InnerProduct(a Vector) E4 { return res } -func (vector Vector) MulByElement(a Vector, b fr.Vector) { +func (vector VectorE4) MulByElement(a VectorE4, b fr.Vector) { N := len(vector) if len(a) != N || len(b) != N { panic("vector.MulByElement: vectors don't have the same length") @@ -211,7 +216,7 @@ func (vector Vector) MulByElement(a Vector, b fr.Vector) { // Butterfly computes the in-place butterfly operation on two vectors of E4 elements // If other overlaps with vector, result is undefined, caller should use a temp vector. -func (vector Vector) Butterfly(other Vector) { +func (vector VectorE4) Butterfly(other VectorE4) { N := len(other) if N != len(vector) { panic("vector.Butterfly: vectors don't have the same length") @@ -231,7 +236,7 @@ func (vector Vector) Butterfly(other Vector) { // ButterflyPair computes the in-place butterfly operation of each pair in the vector // vector[0], vector[1]; vector[2], vector[3]; ... -func (vector Vector) ButterflyPair() { +func (vector VectorE4) ButterflyPair() { N := len(vector) if N%2 != 0 { panic("vector.ButterflyPair: vector length must be even") @@ -253,7 +258,7 @@ func (vector Vector) ButterflyPair() { } } -func (vector Vector) ScalarMulByElement(a Vector, b *fr.Element) { +func (vector VectorE4) ScalarMulByElement(a VectorE4, b *fr.Element) { if len(a) != len(vector) { panic("vector.ScalarMulByElement: vectors don't have the same length") } @@ -270,7 +275,7 @@ func (vector Vector) ScalarMulByElement(a Vector, b *fr.Element) { } // Exp sets vector[i] = a[i]ᵏ for all i -func (vector Vector) Exp(a Vector, k int64) { +func (vector VectorE4) Exp(a VectorE4, k int64) { N := len(a) if N != len(vector) { panic("vector.Exp: vectors don't have the same length") @@ -292,7 +297,7 @@ func (vector Vector) Exp(a Vector, k int64) { v0 := &vector[0] // #nosec G602 we check that N > 0 above a0 := &a[0] // #nosec G602 we check that N > 0 above if v0 == a0 { - base = make(Vector, N) + base = make(VectorE4, N) copy(base, a) } } @@ -310,7 +315,7 @@ func (vector Vector) Exp(a Vector, k int64) { // MulAccByElement multiplies each element of the vector v by the E4 element alpha, // accumulating the result in the same vector. -func (vector Vector) MulAccByElement(scale []fr.Element, alpha *E4) { +func (vector VectorE4) MulAccByElement(scale []fr.Element, alpha *E4) { N := len(vector) if N != len(scale) { panic("MulAccByElement: len(vector) != len(scale)") @@ -324,28 +329,28 @@ func (vector Vector) MulAccByElement(scale []fr.Element, alpha *E4) { } // Equal checks whether two vectors are equal -func (vector Vector) Equal(other Vector) bool { +func (vector VectorE4) Equal(other VectorE4) bool { return slices.Equal(vector, other) } // Len is the number of elements in the collection. -func (vector Vector) Len() int { +func (vector VectorE4) Len() int { return len(vector) } // Less reports whether the element with // index i should sort before the element with index j. -func (vector Vector) Less(i, j int) bool { +func (vector VectorE4) Less(i, j int) bool { return vector[i].Cmp(&vector[j]) == -1 } // Swap swaps the elements with indexes i and j. -func (vector Vector) Swap(i, j int) { +func (vector VectorE4) Swap(i, j int) { vector[i], vector[j] = vector[j], vector[i] } // String implements fmt.Stringer interface -func (vector Vector) String() string { +func (vector VectorE4) String() string { var sbb strings.Builder sbb.Grow(len(vector) * 16) sbb.WriteByte('[') @@ -360,7 +365,7 @@ func (vector Vector) String() string { } // MarshalBinary implements encoding.BinaryMarshaler -func (vector *Vector) MarshalBinary() (data []byte, err error) { +func (vector *VectorE4) MarshalBinary() (data []byte, err error) { var buf bytes.Buffer if _, err = vector.WriteTo(&buf); err != nil { @@ -370,7 +375,7 @@ func (vector *Vector) MarshalBinary() (data []byte, err error) { } // UnmarshalBinary implements encoding.BinaryUnmarshaler -func (vector *Vector) UnmarshalBinary(data []byte) error { +func (vector *VectorE4) UnmarshalBinary(data []byte) error { r := bytes.NewReader(data) _, err := vector.ReadFrom(r) return err @@ -378,7 +383,7 @@ func (vector *Vector) UnmarshalBinary(data []byte) error { // WriteTo implements io.WriterTo and writes a vector of big endian encoded Element. // Length of the vector is encoded as a uint32 on the first 4 bytes. -func (vector *Vector) WriteTo(w io.Writer) (int64, error) { +func (vector *VectorE4) WriteTo(w io.Writer) (int64, error) { // encode slice length if err := binary.Write(w, binary.BigEndian, uint32(len(*vector))); err != nil { @@ -406,7 +411,7 @@ func (vector *Vector) WriteTo(w io.Writer) (int64, error) { return n, nil } -// AsyncReadFrom implements an asynchronous version of [Vector.ReadFrom]. It +// AsyncReadFrom implements an asynchronous version of [VectorE4.ReadFrom]. It // reads the reader r in full and then performs the validation and conversion to // Montgomery form separately in a goroutine. Any error encountered during // reading is returned directly, while errors encountered during @@ -430,7 +435,7 @@ func (vector *Vector) WriteTo(w io.Writer) (int64, error) { // - first 4 bytes: length of the vector as a big-endian uint32 // - for each element of the vector, `4 * fr.Bytes` bytes representing the // element in big-endian encoding. -func (vector *Vector) AsyncReadFrom(r io.Reader) (int64, error, chan error) { // nolint ST1008 +func (vector *VectorE4) AsyncReadFrom(r io.Reader) (int64, error, chan error) { // nolint ST1008 chErr := make(chan error, 1) var bufSizeSlice [4]byte if read, err := io.ReadFull(r, bufSizeSlice[:]); err != nil { @@ -459,7 +464,7 @@ func (vector *Vector) AsyncReadFrom(r io.Reader) (int64, error, chan error) { // *vector = (*vector)[:0] if headerSliceLen == 0 { if *vector == nil { - *vector = Vector{} + *vector = VectorE4{} } close(chErr) return totalRead, nil, chErr @@ -467,7 +472,7 @@ func (vector *Vector) AsyncReadFrom(r io.Reader) (int64, error, chan error) { // for i := uint64(0); i < headerSliceLen; i += maxAllocateSliceLength { if len(*vector) <= int(i) { - *vector = append(*vector, make(Vector, int(min(headerSliceLen-i, maxAllocateSliceLength)))...) + *vector = append(*vector, make(VectorE4, int(min(headerSliceLen-i, maxAllocateSliceLength)))...) } bSlice := unsafe.Slice((*byte)(unsafe.Pointer(&(*vector)[i])), int(min(headerSliceLen-i, maxAllocateSliceLength))*BytesE4) read, err := io.ReadFull(r, bSlice) @@ -540,7 +545,7 @@ func (vector *Vector) AsyncReadFrom(r io.Reader) (int64, error, chan error) { // // - for each element of the vector, `4 * fr.Bytes` bytes representing the element in big-endian encoding. // // The method implements [io.ReaderFrom] interface. -func (vector *Vector) ReadFrom(r io.Reader) (int64, error) { +func (vector *VectorE4) ReadFrom(r io.Reader) (int64, error) { // call the async version and wait for the channel to be closed n, err, chErr := vector.AsyncReadFrom(r) @@ -554,7 +559,7 @@ func (vector *Vector) ReadFrom(r io.Reader) (int64, error) { // i.e. vector[i] = vector[0] * vector[1] * ... * vector[i] // If nbTasks > 1, it uses nbTasks goroutines to compute the prefix product in parallel. // If nbTasks is not provided, it uses the number of CPU cores. -func (vector Vector) PrefixProduct(nbTasks ...int) { +func (vector VectorE4) PrefixProduct(nbTasks ...int) { N := len(vector) if N < 2 { return @@ -614,34 +619,34 @@ func (vector Vector) PrefixProduct(nbTasks ...int) { } -func (vector Vector) prefixProductGeneric() { +func (vector VectorE4) prefixProductGeneric() { for i := 1; i < len(vector); i++ { vector[i].Mul(&vector[i], &vector[i-1]) } } -func vectorAddGeneric(res, a, b Vector) { +func vectorAddGeneric(res, a, b VectorE4) { for i := range len(res) { res[i].Add(&a[i], &b[i]) } } -func vectorSubGeneric(res, a, b Vector) { +func vectorSubGeneric(res, a, b VectorE4) { for i := range len(res) { res[i].Sub(&a[i], &b[i]) } } -func vectorMulGeneric(res, a, b Vector) { +func vectorMulGeneric(res, a, b VectorE4) { for i := range len(res) { res[i].Mul(&a[i], &b[i]) } } -func vectorScalarMulGeneric(res, a Vector, b *E4) { +func vectorScalarMulGeneric(res, a VectorE4, b *E4) { for i := range len(res) { res[i].Mul(&a[i], b) } } -func vectorInnerProductGeneric(a, b Vector) E4 { +func vectorInnerProductGeneric(a, b VectorE4) E4 { var res, tmp E4 for i := range len(a) { tmp.Mul(&a[i], &b[i]) @@ -650,7 +655,7 @@ func vectorInnerProductGeneric(a, b Vector) E4 { return res } -func vectorInnerProductByElementGeneric(a Vector, b fr.Vector) E4 { +func vectorInnerProductByElementGeneric(a VectorE4, b fr.Vector) E4 { var res, tmp E4 for i := range len(a) { tmp.MulByElement(&a[i], &b[i]) @@ -659,7 +664,7 @@ func vectorInnerProductByElementGeneric(a Vector, b fr.Vector) E4 { return res } -func vectorSumGeneric(v Vector) E4 { +func vectorSumGeneric(v VectorE4) E4 { var sum E4 for i := range len(v) { sum.Add(&sum, &v[i]) @@ -667,7 +672,7 @@ func vectorSumGeneric(v Vector) E4 { return sum } -func vectorMulAccByElementGeneric(v Vector, scale []fr.Element, alpha *E4) { +func vectorMulAccByElementGeneric(v VectorE4, scale []fr.Element, alpha *E4) { var tmp E4 for i := range len(v) { tmp.MulByElement(alpha, &scale[i]) @@ -675,13 +680,13 @@ func vectorMulAccByElementGeneric(v Vector, scale []fr.Element, alpha *E4) { } } -func vectorMulByElementGeneric(res, a Vector, b fr.Vector) { +func vectorMulByElementGeneric(res, a VectorE4, b fr.Vector) { for i := range len(res) { res[i].MulByElement(&a[i], &b[i]) } } -func vectorButterflyGeneric(a, b Vector) { +func vectorButterflyGeneric(a, b VectorE4) { for i := range len(a) { Butterfly(&a[i], &b[i]) } diff --git a/field/koalabear/fft/fftext.go b/field/koalabear/fft/fftext.go index 6ab1fab248..e69301d8e5 100644 --- a/field/koalabear/fft/fftext.go +++ b/field/koalabear/fft/fftext.go @@ -57,7 +57,7 @@ func (domain *Domain) FFTExt(a []fext.E4, decimation Decimation, opts ...Option) } } parallel.ExecuteAligned(len(a), 4, func(start, end int) { - va := fext.Vector(a[start:end]) + va := fext.VectorE4(a[start:end]) va.MulByElement(va, cosetTable[start:end]) }, opt.nbTasks) } @@ -126,7 +126,7 @@ func (domain *Domain) FFTInverseExt(a []fext.E4, decimation Decimation, opts ... // scale by CardinalityInv if !opt.coset { parallel.ExecuteAligned(len(a), 4, func(start, end int) { - va := fext.Vector(a[start:end]) + va := fext.VectorE4(a[start:end]) va.ScalarMulByElement(va, &domain.CardinalityInv) }, opt.nbTasks) return @@ -156,7 +156,7 @@ func (domain *Domain) FFTInverseExt(a []fext.E4, decimation Decimation, opts ... } } parallel.ExecuteAligned(len(a), 4, func(start, end int) { - va := fext.Vector(a[start:end]) + va := fext.VectorE4(a[start:end]) va.MulByElement(va, cosetTableInv[start:end]) va.ScalarMulByElement(va, &domain.CardinalityInv) }, opt.nbTasks) @@ -226,8 +226,8 @@ func difFFTExt(a []fext.E4, w koalabear.Element, twiddles [][]koalabear.Element, } func innerDIFWithTwiddlesExt(a []fext.E4, twiddles []koalabear.Element, start, end, m int) { - va0 := fext.Vector(a[start:end]) - va1 := fext.Vector(a[start+m : end+m]) + va0 := fext.VectorE4(a[start:end]) + va1 := fext.VectorE4(a[start+m : end+m]) va0.Butterfly(va1) va1.MulByElement(va1, twiddles[start:end]) } @@ -308,8 +308,8 @@ func ditFFTExt(a []fext.E4, w koalabear.Element, twiddles [][]koalabear.Element, } func innerDITWithTwiddlesExt(a []fext.E4, twiddles []koalabear.Element, start, end, m int) { - va0 := fext.Vector(a[start:end]) - va1 := fext.Vector(a[start+m : end+m]) + va0 := fext.VectorE4(a[start:end]) + va1 := fext.VectorE4(a[start+m : end+m]) va1.MulByElement(va1, twiddles[start:end]) va0.Butterfly(va1) } @@ -357,14 +357,14 @@ func kerDIFNP_512Ext(a []fext.E4, twiddles [][]koalabear.Element, stage int) { for offset := 0; offset < 512; offset += 4 { innerDIFWithTwiddlesExt(a[offset:offset+4], twiddles[stage+7], 0, 2, 2) } - va := fext.Vector(a[:512]) + va := fext.VectorE4(a[:512]) va.ButterflyPair() } func kerDITNP_512Ext(a []fext.E4, twiddles [][]koalabear.Element, stage int) { // code unrolled & generated by internal/generator/fft/template/fftext.go.tmpl - va := fext.Vector(a[:512]) + va := fext.VectorE4(a[:512]) va.ButterflyPair() for offset := 0; offset < 512; offset += 4 { innerDITWithTwiddlesExtM2(a[offset:offset+4], twiddles[stage+7]) diff --git a/field/koalabear/vortex/batch_poly.go b/field/koalabear/vortex/batch_poly.go index 46744e8fb4..1386e3b424 100644 --- a/field/koalabear/vortex/batch_poly.go +++ b/field/koalabear/vortex/batch_poly.go @@ -35,7 +35,7 @@ func BatchEvalFextPolyLagrange(polys [][]fext.E4, x fext.E4, oncoset ...bool) ([ results := make([]fext.E4, len(polys)) parallel.Execute(len(polys), func(start, stop int) { for k := start; k < stop; k++ { - res := fext.Vector(polys[k]).InnerProduct(fext.Vector(lagrangeBasis)) + res := fext.VectorE4(polys[k]).InnerProduct(fext.VectorE4(lagrangeBasis)) results[k] = res } }) @@ -67,7 +67,7 @@ func BatchEvalBasePolyLagrange(polys [][]koalabear.Element, x fext.E4, oncoset . results := make([]fext.E4, len(polys)) parallel.Execute(len(polys), func(start, stop int) { for k := start; k < stop; k++ { - res := fext.Vector(lagrangeBasis).InnerProductByElement(polys[k]) + res := fext.VectorE4(lagrangeBasis).InnerProductByElement(polys[k]) results[k] = res } }) @@ -100,7 +100,7 @@ func ComputeLagrangeBasisAtX(n int, x fext.E4, oncoset ...bool) ([]fext.E4, erro numerator.Inverse(&numerator) // compute x-1, x/ω-1, x/ω²-1, ... - res := make(fext.Vector, n) + res := make(fext.VectorE4, n) res[0] = x for i := 1; i < n; i++ { res[i].MulByElement(&res[i-1], generatorInv) @@ -114,7 +114,7 @@ func ComputeLagrangeBasisAtX(n int, x fext.E4, oncoset ...bool) ([]fext.E4, erro } } if isRootOfUnity != -1 { - res = make(fext.Vector, n) + res = make(fext.VectorE4, n) res[isRootOfUnity].SetOne() return res, nil } diff --git a/field/koalabear/vortex/prover.go b/field/koalabear/vortex/prover.go index a786b09e62..92e768d3f0 100644 --- a/field/koalabear/vortex/prover.go +++ b/field/koalabear/vortex/prover.go @@ -134,7 +134,7 @@ func (ps *ProverState) OpenLinComb(alpha fext.E4) { _ualpha := make([]fext.E4, ps.Params.SizeCodeWord()) var lock sync.Mutex parallel.Execute(nbCodewords, func(start, end int) { - ualpha := make(fext.Vector, ps.Params.SizeCodeWord()) + ualpha := make(fext.VectorE4, ps.Params.SizeCodeWord()) alphaPow := new(fext.E4).SetOne() alphaPow.Exp(alpha, big.NewInt(int64(start))) for i := start; i < end; i++ { diff --git a/field/mamabear/extensions/e3.go b/field/mamabear/extensions/e3.go index 490342cf88..61c604abb5 100644 --- a/field/mamabear/extensions/e3.go +++ b/field/mamabear/extensions/e3.go @@ -43,6 +43,27 @@ func (z *E3) SetOne() *E3 { return z } +// SetInt64 sets z to v (embedded via the unique ring homomorphism ℤ→𝔽) and returns z +func (z *E3) SetInt64(v int64) *E3 { + *z = E3{} + z.A0.SetInt64(v) + return z +} + +// SetUint64 sets z to v (embedded via the unique ring homomorphism ℤ→𝔽) and returns z +func (z *E3) SetUint64(v uint64) *E3 { + *z = E3{} + z.A0.SetUint64(v) + return z +} + +// Div sets z to x / y and returns z +func (z *E3) Div(x, y *E3) *E3 { + var r E3 + r.Inverse(y).Mul(x, &r) + return z.Set(&r) +} + // SetBigInt sets z to v (embedded via the unique ring homomorphism ℤ→𝔽, reduced modulo the base field order) and returns z func (z *E3) SetBigInt(v *big.Int) *E3 { *z = E3{} diff --git a/field/mamabear/extensions/e3_test.go b/field/mamabear/extensions/e3_test.go index 7c381f14d0..eda81b7e86 100644 --- a/field/mamabear/extensions/e3_test.go +++ b/field/mamabear/extensions/e3_test.go @@ -256,18 +256,18 @@ func TestBatchInvertE3(t *testing.T) { } } -// ---- Vector tests ------------------------------------------------------------ +// ---- VectorE3 tests ------------------------------------------------------------ func TestE3VectorButterfly(t *testing.T) { const n = 64 - a := make(Vector, n) - b := make(Vector, n) + a := make(VectorE3, n) + b := make(VectorE3, n) for i := range n { a[i].MustSetRandom() b[i].MustSetRandom() } - aOrig := make(Vector, n) - bOrig := make(Vector, n) + aOrig := make(VectorE3, n) + bOrig := make(VectorE3, n) copy(aOrig, a) copy(bOrig, b) @@ -288,11 +288,11 @@ func TestE3VectorButterfly(t *testing.T) { func TestE3VectorButterflyPair(t *testing.T) { const n = 64 - v := make(Vector, n) + v := make(VectorE3, n) for i := range n { v[i].MustSetRandom() } - orig := make(Vector, n) + orig := make(VectorE3, n) copy(orig, v) v.ButterflyPair() @@ -312,14 +312,14 @@ func TestE3VectorButterflyPair(t *testing.T) { func TestE3VectorMulByElement(t *testing.T) { const n = 32 - a := make(Vector, n) + a := make(VectorE3, n) var scalars [n]fr.Element for i := range n { a[i].MustSetRandom() scalars[i].MustSetRandom() } - res := make(Vector, n) + res := make(VectorE3, n) res.MulByElement(a, scalars[:]) for i := range n { @@ -333,14 +333,14 @@ func TestE3VectorMulByElement(t *testing.T) { func TestE3VectorScalarMulByElement(t *testing.T) { const n = 32 - a := make(Vector, n) + a := make(VectorE3, n) for i := range n { a[i].MustSetRandom() } var s fr.Element s.MustSetRandom() - res := make(Vector, n) + res := make(VectorE3, n) res.ScalarMulByElement(a, &s) for i := range n { @@ -354,7 +354,7 @@ func TestE3VectorScalarMulByElement(t *testing.T) { func TestE3VectorMulAccByElement(t *testing.T) { for _, n := range []int{0, 1, 7, 8, 16, 63, 64, 65, 100, 1000} { - dst := make(Vector, n) + dst := make(VectorE3, n) scale := make([]fr.Element, n) for i := range n { dst[i].MustSetRandom() @@ -362,7 +362,7 @@ func TestE3VectorMulAccByElement(t *testing.T) { } alpha := randE3(t) - expected := make(Vector, n) + expected := make(VectorE3, n) copy(expected, dst) var tmp E3 for i := range n { @@ -432,7 +432,7 @@ func BenchmarkE3Exp(b *testing.B) { func BenchmarkE3VectorMulAccByElement(b *testing.B) { const n = 1 << 16 - dst := make(Vector, n) + dst := make(VectorE3, n) scale := make([]fr.Element, n) for i := range n { dst[i].MustSetRandom() @@ -478,3 +478,53 @@ func TestE3MarshalSetBytesRoundTrip(t *testing.T) { } } } + +func TestE3SetInt64SetUint64(t *testing.T) { + for _, v := range []int64{0, 1, -1, 7, -12345, 1 << 40, -(1 << 40)} { + var z E3 + z.SetInt64(v) + require.Truef(t, z.A1.IsZero(), "coordinate A1 is non-zero for %d", v) + require.Truef(t, z.A2.IsZero(), "coordinate A2 is non-zero for %d", v) + var want fr.Element + want.SetInt64(v) + require.Truef(t, z.A0.Equal(&want), "SetInt64(%d)", v) + + // SetInt64 and SetBigInt agree + var zb E3 + zb.SetBigInt(big.NewInt(v)) + require.Truef(t, z.Equal(&zb), "SetInt64(%d) != SetBigInt(%d)", v, v) + } + for _, v := range []uint64{0, 1, 7, 12345, 1 << 40, 1<<64 - 1} { + var z E3 + z.SetUint64(v) + require.Truef(t, z.A1.IsZero(), "coordinate A1 is non-zero for %d", v) + require.Truef(t, z.A2.IsZero(), "coordinate A2 is non-zero for %d", v) + var want fr.Element + want.SetUint64(v) + require.Truef(t, z.A0.Equal(&want), "SetUint64(%d)", v) + + var zb E3 + zb.SetBigInt(new(big.Int).SetUint64(v)) + require.Truef(t, z.Equal(&zb), "SetUint64(%d) != SetBigInt(%d)", v, v) + } +} + +func TestE3Div(t *testing.T) { + for range 100 { + var x, y, q, back E3 + x.MustSetRandom() + y.MustSetRandom() + if y.IsZero() { + continue + } + q.Div(&x, &y) + back.Mul(&q, &y) + require.True(t, back.Equal(&x), "(x/y)*y != x") + + // aliasing the receiver with the numerator + var xc E3 + xc.Set(&x) + xc.Div(&xc, &y) + require.True(t, xc.Equal(&q), "Div must be alias-safe") + } +} diff --git a/field/mamabear/extensions/e3_vector.go b/field/mamabear/extensions/e3_vector.go index ace2731cce..51a53ef9e0 100644 --- a/field/mamabear/extensions/e3_vector.go +++ b/field/mamabear/extensions/e3_vector.go @@ -19,14 +19,19 @@ func Butterfly(a, b *E3) { b.Sub(&t, b) } +// VectorE3 is a slice of E3 elements. +type VectorE3 []E3 + // Vector is a slice of E3 elements. -type Vector []E3 +// +// Deprecated: use VectorE3. +type Vector = VectorE3 // Butterfly computes the in-place butterfly between two same-length vectors: // // vector[i] = vector[i] + other[i] // other[i] = old vector[i] - other[i] -func (vector Vector) Butterfly(other Vector) { +func (vector VectorE3) Butterfly(other VectorE3) { if len(vector) != len(other) { panic("vector.Butterfly: length mismatch") } @@ -38,7 +43,7 @@ func (vector Vector) Butterfly(other Vector) { // ButterflyPair applies Butterfly to each adjacent pair in the vector: // (vector[0], vector[1]), (vector[2], vector[3]), ... // Length must be even. -func (vector Vector) ButterflyPair() { +func (vector VectorE3) ButterflyPair() { if len(vector)%2 != 0 { panic("vector.ButterflyPair: length must be even") } @@ -48,7 +53,7 @@ func (vector Vector) ButterflyPair() { } // MulByElement sets vector[i] = a[i] * b[i] where b[i] ∈ F_p. -func (vector Vector) MulByElement(a Vector, b fr.Vector) { +func (vector VectorE3) MulByElement(a VectorE3, b fr.Vector) { if len(vector) != len(a) || len(vector) != len(b) { panic("vector.MulByElement: length mismatch") } @@ -61,7 +66,7 @@ func (vector Vector) MulByElement(a Vector, b fr.Vector) { // // Reinterprets the E3 slice as a flat fr.Vector (3 contiguous fr.Elements per E3) // to leverage the optimized fr.Vector.ScalarMul path. -func (vector Vector) ScalarMulByElement(a Vector, b *fr.Element) { +func (vector VectorE3) ScalarMulByElement(a VectorE3, b *fr.Element) { if len(vector) != len(a) { panic("vector.ScalarMulByElement: length mismatch") } @@ -76,7 +81,7 @@ func (vector Vector) ScalarMulByElement(a Vector, b *fr.Element) { } // ScalarMul sets vector[i] = a[i] * b for all i, where b ∈ E3. -func (vector Vector) ScalarMul(a Vector, b *E3) { +func (vector VectorE3) ScalarMul(a VectorE3, b *E3) { if len(vector) != len(a) { panic("vector.ScalarMul: length mismatch") } @@ -86,7 +91,7 @@ func (vector Vector) ScalarMul(a Vector, b *E3) { } // InnerProduct returns ∑ vector[i] * a[i] over E3. -func (vector Vector) InnerProduct(a Vector) E3 { +func (vector VectorE3) InnerProduct(a VectorE3) E3 { if len(vector) != len(a) { panic("vector.InnerProduct: vectors don't have the same length") } @@ -99,7 +104,7 @@ func (vector Vector) InnerProduct(a Vector) E3 { } // InnerProductByElement returns ∑ vector[i] * a[i] where a[i] ∈ F_p. -func (vector Vector) InnerProductByElement(a fr.Vector) E3 { +func (vector VectorE3) InnerProductByElement(a fr.Vector) E3 { if len(vector) != len(a) { panic("vector.InnerProductByElement: vectors don't have the same length") } @@ -124,7 +129,7 @@ const mulAccByElementThreshold = 64 // kernels. Without AVX-512IFMA, fr.Vector.Mul/Add fall back to the same // per-element scalar loop this batching wraps around, so the extra tiling // work is pure overhead — in that case we use the plain scalar loop instead. -func (vector Vector) MulAccByElement(scale []fr.Element, alpha *E3) { +func (vector VectorE3) MulAccByElement(scale []fr.Element, alpha *E3) { n := len(vector) if n != len(scale) { panic("vector.MulAccByElement: length mismatch") @@ -144,7 +149,7 @@ func (vector Vector) MulAccByElement(scale []fr.Element, alpha *E3) { // processing it with fr.Vector's Mul/Add, which dispatch to the AVX-512IFMA // kernels on amd64. Split out from MulAccByElement so its correctness can be // tested independently of cpu.SupportAVX512IFMA. -func mulAccByElementBatched(vector Vector, scale []fr.Element, alpha *E3) { +func mulAccByElementBatched(vector VectorE3, scale []fr.Element, alpha *E3) { n := len(vector) // E3 = {A0, A1, A2} with no padding — safe to reinterpret as 3×fr.Element. @@ -167,7 +172,7 @@ func mulAccByElementBatched(vector Vector, scale []fr.Element, alpha *E3) { flatVector.Add(flatVector, tmp) } -func mulAccByElementGeneric(vector Vector, scale []fr.Element, alpha *E3) { +func mulAccByElementGeneric(vector VectorE3, scale []fr.Element, alpha *E3) { var tmp E3 for i := range vector { tmp.MulByElement(alpha, &scale[i]) diff --git a/field/mamabear/extensions/e3_vector_internal_test.go b/field/mamabear/extensions/e3_vector_internal_test.go index 184e2b36c8..38cda577b9 100644 --- a/field/mamabear/extensions/e3_vector_internal_test.go +++ b/field/mamabear/extensions/e3_vector_internal_test.go @@ -17,7 +17,7 @@ import ( // exercises mulAccByElementBatched. func TestMulAccByElementBatchedMatchesGeneric(t *testing.T) { for _, n := range []int{1, 7, 8, 64, 65, 100, 1000} { - dst := make(Vector, n) + dst := make(VectorE3, n) scale := make([]fr.Element, n) for i := range n { dst[i].MustSetRandom() @@ -26,11 +26,11 @@ func TestMulAccByElementBatchedMatchesGeneric(t *testing.T) { var alpha E3 alpha.MustSetRandom() - expected := make(Vector, n) + expected := make(VectorE3, n) copy(expected, dst) mulAccByElementGeneric(expected, scale, &alpha) - got := make(Vector, n) + got := make(VectorE3, n) copy(got, dst) mulAccByElementBatched(got, scale, &alpha) diff --git a/field/mamabear/fft/fftext.go b/field/mamabear/fft/fftext.go index 77ce282b7a..8592d10b9e 100644 --- a/field/mamabear/fft/fftext.go +++ b/field/mamabear/fft/fftext.go @@ -57,7 +57,7 @@ func (domain *Domain) FFTExt(a []fext.E3, decimation Decimation, opts ...Option) } } parallel.ExecuteAligned(len(a), 1, func(start, end int) { - va := fext.Vector(a[start:end]) + va := fext.VectorE3(a[start:end]) va.MulByElement(va, cosetTable[start:end]) }, opt.nbTasks) } @@ -126,7 +126,7 @@ func (domain *Domain) FFTInverseExt(a []fext.E3, decimation Decimation, opts ... // scale by CardinalityInv if !opt.coset { parallel.ExecuteAligned(len(a), 1, func(start, end int) { - va := fext.Vector(a[start:end]) + va := fext.VectorE3(a[start:end]) va.ScalarMulByElement(va, &domain.CardinalityInv) }, opt.nbTasks) return @@ -156,7 +156,7 @@ func (domain *Domain) FFTInverseExt(a []fext.E3, decimation Decimation, opts ... } } parallel.ExecuteAligned(len(a), 1, func(start, end int) { - va := fext.Vector(a[start:end]) + va := fext.VectorE3(a[start:end]) va.MulByElement(va, cosetTableInv[start:end]) va.ScalarMulByElement(va, &domain.CardinalityInv) }, opt.nbTasks) @@ -226,8 +226,8 @@ func difFFTExt(a []fext.E3, w mamabear.Element, twiddles [][]mamabear.Element, t } func innerDIFWithTwiddlesExt(a []fext.E3, twiddles []mamabear.Element, start, end, m int) { - va0 := fext.Vector(a[start:end]) - va1 := fext.Vector(a[start+m : end+m]) + va0 := fext.VectorE3(a[start:end]) + va1 := fext.VectorE3(a[start+m : end+m]) va0.Butterfly(va1) va1.MulByElement(va1, twiddles[start:end]) } @@ -308,8 +308,8 @@ func ditFFTExt(a []fext.E3, w mamabear.Element, twiddles [][]mamabear.Element, t } func innerDITWithTwiddlesExt(a []fext.E3, twiddles []mamabear.Element, start, end, m int) { - va0 := fext.Vector(a[start:end]) - va1 := fext.Vector(a[start+m : end+m]) + va0 := fext.VectorE3(a[start:end]) + va1 := fext.VectorE3(a[start+m : end+m]) va1.MulByElement(va1, twiddles[start:end]) va0.Butterfly(va1) } @@ -357,14 +357,14 @@ func kerDIFNP_512Ext(a []fext.E3, twiddles [][]mamabear.Element, stage int) { for offset := 0; offset < 512; offset += 4 { innerDIFWithTwiddlesExt(a[offset:offset+4], twiddles[stage+7], 0, 2, 2) } - va := fext.Vector(a[:512]) + va := fext.VectorE3(a[:512]) va.ButterflyPair() } func kerDITNP_512Ext(a []fext.E3, twiddles [][]mamabear.Element, stage int) { // code unrolled & generated by internal/generator/fft/template/fftext.go.tmpl - va := fext.Vector(a[:512]) + va := fext.VectorE3(a[:512]) va.ButterflyPair() for offset := 0; offset < 512; offset += 4 { innerDITWithTwiddlesExtM2(a[offset:offset+4], twiddles[stage+7]) diff --git a/field/mamabear/vortex/batch_poly.go b/field/mamabear/vortex/batch_poly.go index c02f69e78c..3304c0ee6b 100644 --- a/field/mamabear/vortex/batch_poly.go +++ b/field/mamabear/vortex/batch_poly.go @@ -31,7 +31,7 @@ func BatchEvalFextPolyLagrange(polys [][]fext.E3, x fext.E3, oncoset ...bool) ([ results := make([]fext.E3, len(polys)) parallel.Execute(len(polys), func(start, stop int) { for k := start; k < stop; k++ { - res := fext.Vector(polys[k]).InnerProduct(fext.Vector(lagrangeBasis)) + res := fext.VectorE3(polys[k]).InnerProduct(fext.VectorE3(lagrangeBasis)) results[k] = res } }) @@ -60,7 +60,7 @@ func BatchEvalBasePolyLagrange(polys [][]mamabear.Element, x fext.E3, oncoset .. results := make([]fext.E3, len(polys)) parallel.Execute(len(polys), func(start, stop int) { for k := start; k < stop; k++ { - res := fext.Vector(lagrangeBasis).InnerProductByElement(polys[k]) + res := fext.VectorE3(lagrangeBasis).InnerProductByElement(polys[k]) results[k] = res } }) @@ -91,7 +91,7 @@ func ComputeLagrangeBasisAtX(n int, x fext.E3, oncoset ...bool) ([]fext.E3, erro numerator.MulByElement(&numerator, &cardInv) numerator.Inverse(&numerator) - res := make(fext.Vector, n) + res := make(fext.VectorE3, n) res[0] = x for i := 1; i < n; i++ { res[i].MulByElement(&res[i-1], generatorInv) @@ -105,7 +105,7 @@ func ComputeLagrangeBasisAtX(n int, x fext.E3, oncoset ...bool) ([]fext.E3, erro } } if isRootOfUnity != -1 { - res = make(fext.Vector, n) + res = make(fext.VectorE3, n) res[isRootOfUnity].SetOne() return res, nil } diff --git a/field/mamabear/vortex/prover.go b/field/mamabear/vortex/prover.go index 96f1a7f307..3de59867f1 100644 --- a/field/mamabear/vortex/prover.go +++ b/field/mamabear/vortex/prover.go @@ -101,7 +101,7 @@ func (ps *ProverState) OpenLinComb(alpha fext.E3) { _ualpha := make([]fext.E3, ps.Params.SizeCodeWord()) var lock sync.Mutex parallel.Execute(nbCodewords, func(start, end int) { - ualpha := make(fext.Vector, ps.Params.SizeCodeWord()) + ualpha := make(fext.VectorE3, ps.Params.SizeCodeWord()) alphaPow := new(fext.E3).SetOne() alphaPow.Exp(alpha, big.NewInt(int64(start))) for i := start; i < end; i++ { diff --git a/internal/generator/field/template/extensions/e2.go.tmpl b/internal/generator/field/template/extensions/e2.go.tmpl index 114b4a7e34..a8970fe6cf 100644 --- a/internal/generator/field/template/extensions/e2.go.tmpl +++ b/internal/generator/field/template/extensions/e2.go.tmpl @@ -67,6 +67,20 @@ func (z *E2) SetOne() *E2 { return z } +// SetInt64 sets z to v (embedded via the unique ring homomorphism ℤ→𝔽) and returns z +func (z *E2) SetInt64(v int64) *E2 { + *z = E2{} + z.A0.SetInt64(v) + return z +} + +// SetUint64 sets z to v (embedded via the unique ring homomorphism ℤ→𝔽) and returns z +func (z *E2) SetUint64(v uint64) *E2 { + *z = E2{} + z.A0.SetUint64(v) + return z +} + // SetBigInt sets z to v (embedded via the unique ring homomorphism ℤ→𝔽, reduced modulo the base field order) and returns z func (z *E2) SetBigInt(v *big.Int) *E2 { *z = E2{} diff --git a/internal/generator/field/template/extensions/e2_test.go.tmpl b/internal/generator/field/template/extensions/e2_test.go.tmpl index b802a7778e..89c55f9021 100644 --- a/internal/generator/field/template/extensions/e2_test.go.tmpl +++ b/internal/generator/field/template/extensions/e2_test.go.tmpl @@ -593,3 +593,31 @@ func TestE2MarshalSetBytesRoundTrip(t *testing.T) { } } } + +func TestE2SetInt64SetUint64(t *testing.T) { + for _, v := range []int64{0, 1, -1, 7, -12345, 1 << 40, -(1 << 40)} { + var z E2 + z.SetInt64(v) + require.Truef(t, z.A1.IsZero(), "coordinate A1 is non-zero for %d", v) + var want fr.Element + want.SetInt64(v) + require.Truef(t, z.A0.Equal(&want), "SetInt64(%d)", v) + + // SetInt64 and SetBigInt agree + var zb E2 + zb.SetBigInt(big.NewInt(v)) + require.Truef(t, z.Equal(&zb), "SetInt64(%d) != SetBigInt(%d)", v, v) + } + for _, v := range []uint64{0, 1, 7, 12345, 1 << 40, 1<<64 - 1} { + var z E2 + z.SetUint64(v) + require.Truef(t, z.A1.IsZero(), "coordinate A1 is non-zero for %d", v) + var want fr.Element + want.SetUint64(v) + require.Truef(t, z.A0.Equal(&want), "SetUint64(%d)", v) + + var zb E2 + zb.SetBigInt(new(big.Int).SetUint64(v)) + require.Truef(t, z.Equal(&zb), "SetUint64(%d) != SetBigInt(%d)", v, v) + } +} diff --git a/internal/generator/field/template/extensions/e3.go.tmpl b/internal/generator/field/template/extensions/e3.go.tmpl index 247663c1c7..f07396b0cb 100644 --- a/internal/generator/field/template/extensions/e3.go.tmpl +++ b/internal/generator/field/template/extensions/e3.go.tmpl @@ -35,6 +35,27 @@ func (z *E3) SetOne() *E3 { return z } +// SetInt64 sets z to v (embedded via the unique ring homomorphism ℤ→𝔽) and returns z +func (z *E3) SetInt64(v int64) *E3 { + *z = E3{} + z.A0.SetInt64(v) + return z +} + +// SetUint64 sets z to v (embedded via the unique ring homomorphism ℤ→𝔽) and returns z +func (z *E3) SetUint64(v uint64) *E3 { + *z = E3{} + z.A0.SetUint64(v) + return z +} + +// Div sets z to x / y and returns z +func (z *E3) Div(x, y *E3) *E3 { + var r E3 + r.Inverse(y).Mul(x, &r) + return z.Set(&r) +} + // SetBigInt sets z to v (embedded via the unique ring homomorphism ℤ→𝔽, reduced modulo the base field order) and returns z func (z *E3) SetBigInt(v *big.Int) *E3 { *z = E3{} diff --git a/internal/generator/field/template/extensions/e3_test.go.tmpl b/internal/generator/field/template/extensions/e3_test.go.tmpl index 014eeb51b0..5b0a56113a 100644 --- a/internal/generator/field/template/extensions/e3_test.go.tmpl +++ b/internal/generator/field/template/extensions/e3_test.go.tmpl @@ -247,18 +247,18 @@ func TestBatchInvertE3(t *testing.T) { } } -// ---- Vector tests ------------------------------------------------------------ +// ---- VectorE3 tests ------------------------------------------------------------ func TestE3VectorButterfly(t *testing.T) { const n = 64 - a := make(Vector, n) - b := make(Vector, n) + a := make(VectorE3, n) + b := make(VectorE3, n) for i := range n { a[i].MustSetRandom() b[i].MustSetRandom() } - aOrig := make(Vector, n) - bOrig := make(Vector, n) + aOrig := make(VectorE3, n) + bOrig := make(VectorE3, n) copy(aOrig, a) copy(bOrig, b) @@ -279,11 +279,11 @@ func TestE3VectorButterfly(t *testing.T) { func TestE3VectorButterflyPair(t *testing.T) { const n = 64 - v := make(Vector, n) + v := make(VectorE3, n) for i := range n { v[i].MustSetRandom() } - orig := make(Vector, n) + orig := make(VectorE3, n) copy(orig, v) v.ButterflyPair() @@ -303,14 +303,14 @@ func TestE3VectorButterflyPair(t *testing.T) { func TestE3VectorMulByElement(t *testing.T) { const n = 32 - a := make(Vector, n) + a := make(VectorE3, n) var scalars [n]fr.Element for i := range n { a[i].MustSetRandom() scalars[i].MustSetRandom() } - res := make(Vector, n) + res := make(VectorE3, n) res.MulByElement(a, scalars[:]) for i := range n { @@ -324,14 +324,14 @@ func TestE3VectorMulByElement(t *testing.T) { func TestE3VectorScalarMulByElement(t *testing.T) { const n = 32 - a := make(Vector, n) + a := make(VectorE3, n) for i := range n { a[i].MustSetRandom() } var s fr.Element s.MustSetRandom() - res := make(Vector, n) + res := make(VectorE3, n) res.ScalarMulByElement(a, &s) for i := range n { @@ -345,7 +345,7 @@ func TestE3VectorScalarMulByElement(t *testing.T) { func TestE3VectorMulAccByElement(t *testing.T) { for _, n := range []int{0, 1, 7, 8, 16, 63, 64, 65, 100, 1000} { - dst := make(Vector, n) + dst := make(VectorE3, n) scale := make([]fr.Element, n) for i := range n { dst[i].MustSetRandom() @@ -353,7 +353,7 @@ func TestE3VectorMulAccByElement(t *testing.T) { } alpha := randE3(t) - expected := make(Vector, n) + expected := make(VectorE3, n) copy(expected, dst) var tmp E3 for i := range n { @@ -423,7 +423,7 @@ func BenchmarkE3Exp(b *testing.B) { func BenchmarkE3VectorMulAccByElement(b *testing.B) { const n = 1 << 16 - dst := make(Vector, n) + dst := make(VectorE3, n) scale := make([]fr.Element, n) for i := range n { dst[i].MustSetRandom() @@ -469,3 +469,53 @@ func TestE3MarshalSetBytesRoundTrip(t *testing.T) { } } } + +func TestE3SetInt64SetUint64(t *testing.T) { + for _, v := range []int64{0, 1, -1, 7, -12345, 1 << 40, -(1 << 40)} { + var z E3 + z.SetInt64(v) + require.Truef(t, z.A1.IsZero(), "coordinate A1 is non-zero for %d", v) + require.Truef(t, z.A2.IsZero(), "coordinate A2 is non-zero for %d", v) + var want fr.Element + want.SetInt64(v) + require.Truef(t, z.A0.Equal(&want), "SetInt64(%d)", v) + + // SetInt64 and SetBigInt agree + var zb E3 + zb.SetBigInt(big.NewInt(v)) + require.Truef(t, z.Equal(&zb), "SetInt64(%d) != SetBigInt(%d)", v, v) + } + for _, v := range []uint64{0, 1, 7, 12345, 1 << 40, 1<<64 - 1} { + var z E3 + z.SetUint64(v) + require.Truef(t, z.A1.IsZero(), "coordinate A1 is non-zero for %d", v) + require.Truef(t, z.A2.IsZero(), "coordinate A2 is non-zero for %d", v) + var want fr.Element + want.SetUint64(v) + require.Truef(t, z.A0.Equal(&want), "SetUint64(%d)", v) + + var zb E3 + zb.SetBigInt(new(big.Int).SetUint64(v)) + require.Truef(t, z.Equal(&zb), "SetUint64(%d) != SetBigInt(%d)", v, v) + } +} + +func TestE3Div(t *testing.T) { + for range 100 { + var x, y, q, back E3 + x.MustSetRandom() + y.MustSetRandom() + if y.IsZero() { + continue + } + q.Div(&x, &y) + back.Mul(&q, &y) + require.True(t, back.Equal(&x), "(x/y)*y != x") + + // aliasing the receiver with the numerator + var xc E3 + xc.Set(&x) + xc.Div(&xc, &y) + require.True(t, xc.Equal(&q), "Div must be alias-safe") + } +} diff --git a/internal/generator/field/template/extensions/e3vector.go.tmpl b/internal/generator/field/template/extensions/e3vector.go.tmpl index d12cafe681..10a6c17b25 100644 --- a/internal/generator/field/template/extensions/e3vector.go.tmpl +++ b/internal/generator/field/template/extensions/e3vector.go.tmpl @@ -12,14 +12,19 @@ func Butterfly(a, b *E3) { b.Sub(&t, b) } +// VectorE3 is a slice of E3 elements. +type VectorE3 []E3 + // Vector is a slice of E3 elements. -type Vector []E3 +// +// Deprecated: use VectorE3. +type Vector = VectorE3 // Butterfly computes the in-place butterfly between two same-length vectors: // // vector[i] = vector[i] + other[i] // other[i] = old vector[i] - other[i] -func (vector Vector) Butterfly(other Vector) { +func (vector VectorE3) Butterfly(other VectorE3) { if len(vector) != len(other) { panic("vector.Butterfly: length mismatch") } @@ -31,7 +36,7 @@ func (vector Vector) Butterfly(other Vector) { // ButterflyPair applies Butterfly to each adjacent pair in the vector: // (vector[0], vector[1]), (vector[2], vector[3]), ... // Length must be even. -func (vector Vector) ButterflyPair() { +func (vector VectorE3) ButterflyPair() { if len(vector)%2 != 0 { panic("vector.ButterflyPair: length must be even") } @@ -41,7 +46,7 @@ func (vector Vector) ButterflyPair() { } // MulByElement sets vector[i] = a[i] * b[i] where b[i] ∈ F_p. -func (vector Vector) MulByElement(a Vector, b fr.Vector) { +func (vector VectorE3) MulByElement(a VectorE3, b fr.Vector) { if len(vector) != len(a) || len(vector) != len(b) { panic("vector.MulByElement: length mismatch") } @@ -54,7 +59,7 @@ func (vector Vector) MulByElement(a Vector, b fr.Vector) { // // Reinterprets the E3 slice as a flat fr.Vector (3 contiguous fr.Elements per E3) // to leverage the optimized fr.Vector.ScalarMul path. -func (vector Vector) ScalarMulByElement(a Vector, b *fr.Element) { +func (vector VectorE3) ScalarMulByElement(a VectorE3, b *fr.Element) { if len(vector) != len(a) { panic("vector.ScalarMulByElement: length mismatch") } @@ -69,7 +74,7 @@ func (vector Vector) ScalarMulByElement(a Vector, b *fr.Element) { } // ScalarMul sets vector[i] = a[i] * b for all i, where b ∈ E3. -func (vector Vector) ScalarMul(a Vector, b *E3) { +func (vector VectorE3) ScalarMul(a VectorE3, b *E3) { if len(vector) != len(a) { panic("vector.ScalarMul: length mismatch") } @@ -79,7 +84,7 @@ func (vector Vector) ScalarMul(a Vector, b *E3) { } // InnerProduct returns ∑ vector[i] * a[i] over E3. -func (vector Vector) InnerProduct(a Vector) E3 { +func (vector VectorE3) InnerProduct(a VectorE3) E3 { if len(vector) != len(a) { panic("vector.InnerProduct: vectors don't have the same length") } @@ -92,7 +97,7 @@ func (vector Vector) InnerProduct(a Vector) E3 { } // InnerProductByElement returns ∑ vector[i] * a[i] where a[i] ∈ F_p. -func (vector Vector) InnerProductByElement(a fr.Vector) E3 { +func (vector VectorE3) InnerProductByElement(a fr.Vector) E3 { if len(vector) != len(a) { panic("vector.InnerProductByElement: vectors don't have the same length") } @@ -117,7 +122,7 @@ const mulAccByElementThreshold = 64 // kernels. Without AVX-512IFMA, fr.Vector.Mul/Add fall back to the same // per-element scalar loop this batching wraps around, so the extra tiling // work is pure overhead — in that case we use the plain scalar loop instead. -func (vector Vector) MulAccByElement(scale []fr.Element, alpha *E3) { +func (vector VectorE3) MulAccByElement(scale []fr.Element, alpha *E3) { n := len(vector) if n != len(scale) { panic("vector.MulAccByElement: length mismatch") @@ -137,7 +142,7 @@ func (vector Vector) MulAccByElement(scale []fr.Element, alpha *E3) { // processing it with fr.Vector's Mul/Add, which dispatch to the AVX-512IFMA // kernels on amd64. Split out from MulAccByElement so its correctness can be // tested independently of cpu.SupportAVX512IFMA. -func mulAccByElementBatched(vector Vector, scale []fr.Element, alpha *E3) { +func mulAccByElementBatched(vector VectorE3, scale []fr.Element, alpha *E3) { n := len(vector) // E3 = {A0, A1, A2} with no padding — safe to reinterpret as 3×fr.Element. @@ -160,7 +165,7 @@ func mulAccByElementBatched(vector Vector, scale []fr.Element, alpha *E3) { flatVector.Add(flatVector, tmp) } -func mulAccByElementGeneric(vector Vector, scale []fr.Element, alpha *E3) { +func mulAccByElementGeneric(vector VectorE3, scale []fr.Element, alpha *E3) { var tmp E3 for i := range vector { tmp.MulByElement(alpha, &scale[i]) diff --git a/internal/generator/field/template/extensions/e3vector_test.go.tmpl b/internal/generator/field/template/extensions/e3vector_test.go.tmpl index be282558d6..5abe9dca81 100644 --- a/internal/generator/field/template/extensions/e3vector_test.go.tmpl +++ b/internal/generator/field/template/extensions/e3vector_test.go.tmpl @@ -10,7 +10,7 @@ import ( // exercises mulAccByElementBatched. func TestMulAccByElementBatchedMatchesGeneric(t *testing.T) { for _, n := range []int{1, 7, 8, 64, 65, 100, 1000} { - dst := make(Vector, n) + dst := make(VectorE3, n) scale := make([]fr.Element, n) for i := range n { dst[i].MustSetRandom() @@ -19,11 +19,11 @@ func TestMulAccByElementBatchedMatchesGeneric(t *testing.T) { var alpha E3 alpha.MustSetRandom() - expected := make(Vector, n) + expected := make(VectorE3, n) copy(expected, dst) mulAccByElementGeneric(expected, scale, &alpha) - got := make(Vector, n) + got := make(VectorE3, n) copy(got, dst) mulAccByElementBatched(got, scale, &alpha) diff --git a/internal/generator/field/template/extensions/e4.go.tmpl b/internal/generator/field/template/extensions/e4.go.tmpl index 6af571dfcd..96a77c6bf2 100644 --- a/internal/generator/field/template/extensions/e4.go.tmpl +++ b/internal/generator/field/template/extensions/e4.go.tmpl @@ -78,6 +78,20 @@ func (z *E4) SetOne() *E4 { return z } +// SetInt64 sets z to v (embedded via the unique ring homomorphism ℤ→𝔽) and returns z +func (z *E4) SetInt64(v int64) *E4 { + *z = E4{} + z.B0.A0.SetInt64(v) + return z +} + +// SetUint64 sets z to v (embedded via the unique ring homomorphism ℤ→𝔽) and returns z +func (z *E4) SetUint64(v uint64) *E4 { + *z = E4{} + z.B0.A0.SetUint64(v) + return z +} + // SetBigInt sets z to v (embedded via the unique ring homomorphism ℤ→𝔽, reduced modulo the base field order) and returns z func (z *E4) SetBigInt(v *big.Int) *E4 { *z = E4{} diff --git a/internal/generator/field/template/extensions/e4_test.go.tmpl b/internal/generator/field/template/extensions/e4_test.go.tmpl index 2b6875a1d0..968f4662b4 100644 --- a/internal/generator/field/template/extensions/e4_test.go.tmpl +++ b/internal/generator/field/template/extensions/e4_test.go.tmpl @@ -284,8 +284,8 @@ func TestVectorOps(t *testing.T) { } properties := gopter.NewProperties(parameters) - addVector := func(a, b Vector) bool { - c := make(Vector, len(a)) + addVector := func(a, b VectorE4) bool { + c := make(VectorE4, len(a)) c.Add(a, b) for i := range len(a) { @@ -298,8 +298,8 @@ func TestVectorOps(t *testing.T) { return true } - subVector := func(a, b Vector) bool { - c := make(Vector, len(a)) + subVector := func(a, b VectorE4) bool { + c := make(VectorE4, len(a)) c.Sub(a, b) for i := range len(a) { @@ -312,8 +312,8 @@ func TestVectorOps(t *testing.T) { return true } - scalarMulVector := func(a Vector, b E4) bool { - c := make(Vector, len(a)) + scalarMulVector := func(a VectorE4, b E4) bool { + c := make(VectorE4, len(a)) c.ScalarMul(a, &b) for i := range len(a) { @@ -326,7 +326,7 @@ func TestVectorOps(t *testing.T) { return true } - sumVector := func(a Vector) bool { + sumVector := func(a VectorE4) bool { var sum E4 computed := a.Sum() for i := range len(a) { @@ -336,7 +336,7 @@ func TestVectorOps(t *testing.T) { return sum.Equal(&computed) } - innerProductVector := func(a, b Vector) bool { + innerProductVector := func(a, b VectorE4) bool { computed := a.InnerProduct(b) var innerProduct E4 for i := range len(a) { @@ -348,8 +348,8 @@ func TestVectorOps(t *testing.T) { return innerProduct.Equal(&computed) } - mulVector := func(a, b Vector) bool { - c := make(Vector, len(a)) + mulVector := func(a, b VectorE4) bool { + c := make(VectorE4, len(a)) a[0].B0.A0.SetUint64(0x24) b[0].B0.A0.SetUint64(0x42) c.Mul(a, b) @@ -414,8 +414,8 @@ func TestVectorOps(t *testing.T) { )) properties.Property(fmt.Sprintf("vector scalar multiplication by element %d - %s", size, gp.label), prop.ForAll( - func(a Vector, b fr.Element) bool { - c := make(Vector, len(a)) + func(a VectorE4, b fr.Element) bool { + c := make(VectorE4, len(a)) c.ScalarMulByElement(a, &b) for i := range len(a) { var tmp E4 @@ -431,8 +431,8 @@ func TestVectorOps(t *testing.T) { )) properties.Property(fmt.Sprintf("vector multiplication by element %d - %s", size, gp.label), prop.ForAll( - func(a Vector, b fr.Vector) bool { - c := make(Vector, len(a)) + func(a VectorE4, b fr.Vector) bool { + c := make(VectorE4, len(a)) c.MulByElement(a, b) for i := range len(a) { var tmp E4 @@ -449,12 +449,12 @@ func TestVectorOps(t *testing.T) { // checking that in-place butterfly works as intended; properties.Property(fmt.Sprintf("vector butterfly %d - %s", size, gp.label), prop.ForAll( - func(a, b Vector) bool { + func(a, b VectorE4) bool { if len(a) != len(b) { return false } - c := make(Vector, len(a)) - d := make(Vector, len(a)) + c := make(VectorE4, len(a)) + d := make(VectorE4, len(a)) copy(c, a) copy(d, b) c.Butterfly(d) @@ -476,11 +476,11 @@ func TestVectorOps(t *testing.T) { )) properties.Property(fmt.Sprintf("vector butterfly pair %d - %s", size, gp.label), prop.ForAll( - func(a Vector) bool { + func(a VectorE4) bool { if len(a)%2 != 0 { return true // skip odd-sized vectors } - c := make(Vector, len(a)) + c := make(VectorE4, len(a)) copy(c, a) c.ButterflyPair() for i := 0; i < len(a); i += 2 { @@ -498,7 +498,7 @@ func TestVectorOps(t *testing.T) { )) properties.Property(fmt.Sprintf("vector inner product by element %d - %s", size, gp.label), prop.ForAll( - func(a Vector, b fr.Vector) bool { + func(a VectorE4, b fr.Vector) bool { computed := a.InnerProductByElement(b) var innerProduct E4 for i := range len(a) { @@ -518,23 +518,23 @@ func TestVectorOps(t *testing.T) { properties.TestingRun(t, gopter.NewFormatedReporter(false, 260, os.Stdout)) } -// TestVectorExp tests the Exp method for Vector type. +// TestVectorExp tests the Exp method for VectorE4 type. func TestVectorExp(t *testing.T) { assert := require.New(t) // Test with empty vector - empty := make(Vector, 0) - expEmpty := make(Vector, 0) + empty := make(VectorE4, 0) + expEmpty := make(VectorE4, 0) expEmpty.Exp(empty, 5) assert.Equal(0, len(expEmpty), "Exp of empty vector should be empty") // Test with vector of ones and exponent 0 const size = 32 - v := make(Vector, size) + v := make(VectorE4, size) for i := range v { v[i].SetOne() } - expZero := make(Vector, size) + expZero := make(VectorE4, size) expZero.Exp(v, 0) for i := range expZero { assert.True(expZero[i].IsOne(), "Exp(x, 0) should be one for all elements") @@ -544,14 +544,14 @@ func TestVectorExp(t *testing.T) { for i := range v { v[i].MustSetRandom() } - expOne := make(Vector, size) + expOne := make(VectorE4, size) expOne.Exp(v, 1) for i := range v { assert.True(expOne[i].Equal(&v[i]), "Exp(x, 1) should be x for all elements") } // Test with random vector and exponent 2 - expTwo := make(Vector, size) + expTwo := make(VectorE4, size) expTwo.Exp(v, 2) for i := range v { var sq E4 @@ -561,7 +561,7 @@ func TestVectorExp(t *testing.T) { // Test with random vector and exponent k k := int64(7) - expK := make(Vector, size) + expK := make(VectorE4, size) expK.Exp(v, k) for i := range v { var mul E4 @@ -573,7 +573,7 @@ func TestVectorExp(t *testing.T) { } // Test to check v.Exp(v, k) is correct (no modification of v during the process) - vCopy := make(Vector, size) + vCopy := make(VectorE4, size) copy(vCopy, v) vCopy.Exp(vCopy, k) for i := range v { @@ -581,7 +581,7 @@ func TestVectorExp(t *testing.T) { } // Test with random vector and negative exponent -1 - expNegOne := make(Vector, size) + expNegOne := make(VectorE4, size) expNegOne.Exp(v, -1) for i := range v { var inv E4 @@ -593,7 +593,7 @@ func TestVectorExp(t *testing.T) { } // prefixProductGeneric computes the prefix product of the vector in place (single-threaded). -func prefixProductGeneric(vector Vector) { +func prefixProductGeneric(vector VectorE4) { if len(vector) == 0 { return } @@ -602,8 +602,8 @@ func prefixProductGeneric(vector Vector) { } } -func randomVector(size int) Vector { - v := make(Vector, size) +func randomVector(size int) VectorE4 { + v := make(VectorE4, size) for i := range v { v[i].MustSetRandom() } @@ -612,8 +612,8 @@ func randomVector(size int) Vector { func TestPrefixProduct_EmptyVector(t *testing.T) { assert := require.New(t) - v := make(Vector, 0) - expected := make(Vector, 0) + v := make(VectorE4, 0) + expected := make(VectorE4, 0) prefixProductGeneric(expected) v.PrefixProduct() assert.Equal(expected, v) @@ -626,7 +626,7 @@ func TestPrefixProduct_VariousNbTasks(t *testing.T) { for _, size := range sizes { for _, nbTasks := range nbTasksList { v := randomVector(size) - expected := make(Vector, size) + expected := make(VectorE4, size) copy(expected, v) prefixProductGeneric(expected) v.PrefixProduct(nbTasks) @@ -641,8 +641,8 @@ func TestVectorEmptyOps(t *testing.T) { var sum, inner, scalar E4 scalar.MustSetRandom() - empty := make(Vector, 0) - result := make(Vector, 0) + empty := make(VectorE4, 0) + result := make(VectorE4, 0) assert.NotPanics(func() { result.Add(empty, empty) }) assert.NotPanics(func() { result.Sub(empty, empty) }) @@ -659,12 +659,12 @@ func TestVectorEmptyOps(t *testing.T) { func TestVectorSort(t *testing.T) { assert := require.New(t) - v := make(Vector, 3) + v := make(VectorE4, 3) v[0].B0.A0.SetUint64(2) v[1].B0.A0.SetUint64(3) v[2].B0.A0.SetUint64(1) - expected := make(Vector, 3) + expected := make(VectorE4, 3) expected[0].B0.A0.SetUint64(1) expected[1].B0.A0.SetUint64(2) expected[2].B0.A0.SetUint64(3) @@ -679,7 +679,7 @@ func TestVectorSort(t *testing.T) { func TestVectorRoundTrip(t *testing.T) { assert := require.New(t) - v1 := make(Vector, 3) + v1 := make(VectorE4, 3) v1[0].MustSetRandom() v1[1].MustSetRandom() v1[2].MustSetRandom() @@ -687,7 +687,7 @@ func TestVectorRoundTrip(t *testing.T) { b, err := v1.MarshalBinary() assert.NoError(err) - var v2, v3 Vector + var v2, v3 VectorE4 err = v2.UnmarshalBinary(b) assert.NoError(err) @@ -702,12 +702,12 @@ func TestVectorRoundTrip(t *testing.T) { func TestVectorEmptyRoundTrip(t *testing.T) { assert := require.New(t) - v1 := make(Vector, 0) + v1 := make(VectorE4, 0) b, err := v1.MarshalBinary() assert.NoError(err) - var v2, v3 Vector + var v2, v3 VectorE4 err = v2.UnmarshalBinary(b) assert.NoError(err) @@ -722,7 +722,7 @@ func TestVectorEmptyRoundTrip(t *testing.T) { func TestVectorReuseSliceDeserialization(t *testing.T) { assert := require.New(t) - v1 := make(Vector, 4) + v1 := make(VectorE4, 4) for i := range v1 { v1[i].MustSetRandom() } @@ -730,7 +730,7 @@ func TestVectorReuseSliceDeserialization(t *testing.T) { assert.NoError(err) const capacity = 16 - v2 := make(Vector, capacity) + v2 := make(VectorE4, capacity) n, err, errCh := v2.AsyncReadFrom(bytes.NewReader(buf)) assert.Equal(int64(len(buf)), n) assert.NoError(err) @@ -739,7 +739,7 @@ func TestVectorReuseSliceDeserialization(t *testing.T) { assert.Equal(capacity, cap(v2)) assert.True(reflect.DeepEqual(v1, v2)) - v3 := make(Vector, capacity) + v3 := make(VectorE4, capacity) n, err = v3.ReadFrom(bytes.NewReader(buf)) assert.Equal(int64(len(buf)), n) assert.NoError(err) @@ -755,12 +755,12 @@ func TestVectorReadTamperedHeader(t *testing.T) { var input [4]byte binary.BigEndian.PutUint32(input[:], 1<<12) - newVector := func() Vector { - v := make(Vector, 1) + newVector := func() VectorE4 { + v := make(VectorE4, 1) v[0].SetOne() return v } - assertUnchanged := func(v Vector) { + assertUnchanged := func(v VectorE4) { assert.Len(v, 1) assert.Equal(1, cap(v)) var one E4 @@ -791,7 +791,7 @@ func TestVectorReadTamperedHeader(t *testing.T) { func TestVectorReadTamperedHeaderWithoutLen(t *testing.T) { assert := require.New(t) - v1 := make(Vector, 4) + v1 := make(VectorE4, 4) for i := range v1 { v1[i].MustSetRandom() } @@ -806,12 +806,12 @@ func TestVectorReadTamperedHeaderWithoutLen(t *testing.T) { r := readerWithoutLen() _, hasLen := r.(interface{ Len() int }) assert.False(hasLen) - var v2 Vector + var v2 VectorE4 n, err := v2.ReadFrom(r) assert.Equal(int64(len(buf)), n) assert.Error(err) - var v3 Vector + var v3 VectorE4 n, err, errCh := v3.AsyncReadFrom(readerWithoutLen()) assert.Equal(int64(len(buf)), n) assert.Error(err) @@ -819,7 +819,7 @@ func TestVectorReadTamperedHeaderWithoutLen(t *testing.T) { assert.False(open) } -func (vector *Vector) unmarshalBinaryAsync(data []byte) error { +func (vector *VectorE4) unmarshalBinaryAsync(data []byte) error { r := bytes.NewReader(data) _, err, chErr := vector.AsyncReadFrom(r) if err != nil { @@ -914,9 +914,9 @@ func BenchmarkVectorOps(b *testing.B) { // note; to benchmark against "no asm" version, use the following // build tag: -tags purego const N = 1 << 20 - a1 := make(Vector, N) - b1 := make(Vector, N) - c1 := make(Vector, N) + a1 := make(VectorE4, N) + b1 := make(VectorE4, N) + c1 := make(VectorE4, N) b2 := make(fr.Vector, N) for i := 1; i < N; i++ { a1[i-1].MustSetRandom() @@ -997,7 +997,7 @@ func BenchmarkVectorOps(b *testing.B) { func BenchmarkPrefixProduct(b *testing.B) { const N = 1 << 19 - a1 := make(Vector, N) + a1 := make(VectorE4, N) for i := range N { a1[i].MustSetRandom() } @@ -1021,7 +1021,7 @@ func BenchmarkPrefixProduct(b *testing.B) { func BenchmarkVectorSerialization(b *testing.B) { const N = 1 << 15 - a1 := make(Vector, N) + a1 := make(VectorE4, N) for i := 1; i < N; i++ { a1[i-1].MustSetRandom() } @@ -1041,7 +1041,7 @@ func BenchmarkVectorSerialization(b *testing.B) { } b.Run("UnmarshalBinary", func(b *testing.B) { - var a2 Vector + var a2 VectorE4 b.ResetTimer() for range b.N { err := a2.UnmarshalBinary(data) @@ -1052,7 +1052,7 @@ func BenchmarkVectorSerialization(b *testing.B) { }) b.Run("unmarshalBinaryAsync", func(b *testing.B) { - var a2 Vector + var a2 VectorE4 b.ResetTimer() for range b.N { err := a2.unmarshalBinaryAsync(data) @@ -1065,7 +1065,7 @@ func BenchmarkVectorSerialization(b *testing.B) { func genZeroVector(size int) gopter.Gen { return func(*gopter.GenParameters) *gopter.GenResult { - return gopter.NewGenResult(make(Vector, size), gopter.NoShrinker) + return gopter.NewGenResult(make(VectorE4, size), gopter.NoShrinker) } } @@ -1073,7 +1073,7 @@ func genMaxVector(size int) gopter.Gen { return func(*gopter.GenParameters) *gopter.GenResult { qMinusOne := fr.Element{ {{.Q}} } qMinusOne[0]-- - v := make(Vector, size) + v := make(VectorE4, size) for i := range v { v[i].B0.A0 = qMinusOne v[i].B0.A1 = qMinusOne @@ -1086,7 +1086,7 @@ func genMaxVector(size int) gopter.Gen { func genVector(size int) gopter.Gen { return func(genParams *gopter.GenParameters) *gopter.GenResult { - v := make(Vector, size) + v := make(VectorE4, size) gen := genE4() for i := range v { val, ok := gen(genParams).Retrieve() @@ -1157,3 +1157,35 @@ func TestE4MarshalSetBytesRoundTrip(t *testing.T) { } } } + +func TestE4SetInt64SetUint64(t *testing.T) { + for _, v := range []int64{0, 1, -1, 7, -12345, 1 << 40, -(1 << 40)} { + var z E4 + z.SetInt64(v) + require.Truef(t, z.B0.A1.IsZero(), "coordinate B0.A1 is non-zero for %d", v) + require.Truef(t, z.B1.A0.IsZero(), "coordinate B1.A0 is non-zero for %d", v) + require.Truef(t, z.B1.A1.IsZero(), "coordinate B1.A1 is non-zero for %d", v) + var want fr.Element + want.SetInt64(v) + require.Truef(t, z.B0.A0.Equal(&want), "SetInt64(%d)", v) + + // SetInt64 and SetBigInt agree + var zb E4 + zb.SetBigInt(big.NewInt(v)) + require.Truef(t, z.Equal(&zb), "SetInt64(%d) != SetBigInt(%d)", v, v) + } + for _, v := range []uint64{0, 1, 7, 12345, 1 << 40, 1<<64 - 1} { + var z E4 + z.SetUint64(v) + require.Truef(t, z.B0.A1.IsZero(), "coordinate B0.A1 is non-zero for %d", v) + require.Truef(t, z.B1.A0.IsZero(), "coordinate B1.A0 is non-zero for %d", v) + require.Truef(t, z.B1.A1.IsZero(), "coordinate B1.A1 is non-zero for %d", v) + var want fr.Element + want.SetUint64(v) + require.Truef(t, z.B0.A0.Equal(&want), "SetUint64(%d)", v) + + var zb E4 + zb.SetBigInt(new(big.Int).SetUint64(v)) + require.Truef(t, z.Equal(&zb), "SetUint64(%d) != SetBigInt(%d)", v, v) + } +} diff --git a/internal/generator/field/template/extensions/vector.go.tmpl b/internal/generator/field/template/extensions/vector.go.tmpl index 1f6c7e8708..16629bcd61 100644 --- a/internal/generator/field/template/extensions/vector.go.tmpl +++ b/internal/generator/field/template/extensions/vector.go.tmpl @@ -19,10 +19,15 @@ import ( {{- end}} ) -// Vector represents a vector of E4 elements -type Vector []E4 +// VectorE4 represents a vector of E4 elements +type VectorE4 []E4 -func (vector Vector) Add(a, b Vector) { +// Vector is a slice of E4 elements. +// +// Deprecated: use VectorE4. +type Vector = VectorE4 + +func (vector VectorE4) Add(a, b VectorE4) { N := len(a) if N != len(b) || N != len(vector) { panic("vector.Add: vectors don't have the same length") @@ -44,7 +49,7 @@ func (vector Vector) Add(a, b Vector) { {{- end}} } -func (vector Vector) Sub(a, b Vector) { +func (vector VectorE4) Sub(a, b VectorE4) { N := len(a) if N != len(b) || N != len(vector) { panic("vector.Sub: vectors don't have the same length") @@ -66,7 +71,7 @@ func (vector Vector) Sub(a, b Vector) { {{- end}} } -func (vector Vector) Mul(a, b Vector) { +func (vector VectorE4) Mul(a, b VectorE4) { N := len(a) if N != len(b) || N != len(vector) { panic("vector.Mul: vectors don't have the same length") @@ -88,7 +93,7 @@ func (vector Vector) Mul(a, b Vector) { {{- end}} } -func (vector Vector) ScalarMul(a Vector, b *E4) { +func (vector VectorE4) ScalarMul(a VectorE4, b *E4) { N := len(a) if N != len(vector) { panic("vector.ScalarMul: vectors don't have the same length") @@ -111,7 +116,7 @@ func (vector Vector) ScalarMul(a Vector, b *E4) { } // Sum computes the sum of all elements in the vector. -func (vector Vector) Sum() E4 { +func (vector VectorE4) Sum() E4 { {{- if .IsKoalaBear}} const blockSize = 2 N := len(vector) @@ -140,7 +145,7 @@ func (vector Vector) Sum() E4 { {{- end}} } -func (vector Vector) InnerProductByElement(a fr.Vector) E4 { +func (vector VectorE4) InnerProductByElement(a fr.Vector) E4 { N := len(vector) if len(a) != N { panic("vector.InnerProduct: vectors don't have the same length") @@ -165,7 +170,7 @@ func (vector Vector) InnerProductByElement(a fr.Vector) E4 { {{- end}} } -func (vector Vector) InnerProduct(a Vector) E4 { +func (vector VectorE4) InnerProduct(a VectorE4) E4 { N := len(vector) if len(a) != N { panic("vector.InnerProduct: vectors don't have the same length") @@ -212,7 +217,7 @@ func (vector Vector) InnerProduct(a Vector) E4 { {{- end}} } -func (vector Vector) MulByElement(a Vector, b fr.Vector) { +func (vector VectorE4) MulByElement(a VectorE4, b fr.Vector) { N := len(vector) if len(a) != N || len(b) != N { panic("vector.MulByElement: vectors don't have the same length") @@ -241,7 +246,7 @@ func (vector Vector) MulByElement(a Vector, b fr.Vector) { // Butterfly computes the in-place butterfly operation on two vectors of E4 elements // If other overlaps with vector, result is undefined, caller should use a temp vector. -func (vector Vector) Butterfly(other Vector) { +func (vector VectorE4) Butterfly(other VectorE4) { N := len(other) if N != len(vector) { panic("vector.Butterfly: vectors don't have the same length") @@ -265,7 +270,7 @@ func (vector Vector) Butterfly(other Vector) { // ButterflyPair computes the in-place butterfly operation of each pair in the vector // vector[0], vector[1]; vector[2], vector[3]; ... -func (vector Vector) ButterflyPair() { +func (vector VectorE4) ButterflyPair() { N := len(vector) if N%2 != 0 { panic("vector.ButterflyPair: vector length must be even") @@ -293,7 +298,7 @@ func (vector Vector) ButterflyPair() { {{- end}} } -func (vector Vector) ScalarMulByElement(a Vector, b *fr.Element) { +func (vector VectorE4) ScalarMulByElement(a VectorE4, b *fr.Element) { if len(a) != len(vector) { panic("vector.ScalarMulByElement: vectors don't have the same length") } @@ -310,7 +315,7 @@ func (vector Vector) ScalarMulByElement(a Vector, b *fr.Element) { } // Exp sets vector[i] = a[i]ᵏ for all i -func (vector Vector) Exp(a Vector, k int64) { +func (vector VectorE4) Exp(a VectorE4, k int64) { N := len(a) if N != len(vector) { panic("vector.Exp: vectors don't have the same length") @@ -332,7 +337,7 @@ func (vector Vector) Exp(a Vector, k int64) { v0 := &vector[0] // #nosec G602 we check that N > 0 above a0 := &a[0] // #nosec G602 we check that N > 0 above if v0 == a0 { - base = make(Vector, N) + base = make(VectorE4, N) copy(base, a) } } @@ -350,7 +355,7 @@ func (vector Vector) Exp(a Vector, k int64) { // MulAccByElement multiplies each element of the vector v by the E4 element alpha, // accumulating the result in the same vector. -func (vector Vector) MulAccByElement(scale []fr.Element, alpha *E4) { +func (vector VectorE4) MulAccByElement(scale []fr.Element, alpha *E4) { N := len(vector) if N != len(scale) { panic("MulAccByElement: len(vector) != len(scale)") @@ -369,28 +374,28 @@ func (vector Vector) MulAccByElement(scale []fr.Element, alpha *E4) { // Equal checks whether two vectors are equal -func (vector Vector) Equal(other Vector) bool { +func (vector VectorE4) Equal(other VectorE4) bool { return slices.Equal(vector, other) } // Len is the number of elements in the collection. -func (vector Vector) Len() int { +func (vector VectorE4) Len() int { return len(vector) } // Less reports whether the element with // index i should sort before the element with index j. -func (vector Vector) Less(i, j int) bool { +func (vector VectorE4) Less(i, j int) bool { return vector[i].Cmp(&vector[j]) == -1 } // Swap swaps the elements with indexes i and j. -func (vector Vector) Swap(i, j int) { +func (vector VectorE4) Swap(i, j int) { vector[i], vector[j] = vector[j], vector[i] } // String implements fmt.Stringer interface -func (vector Vector) String() string { +func (vector VectorE4) String() string { var sbb strings.Builder sbb.Grow(len(vector) * 16) sbb.WriteByte('[') @@ -406,7 +411,7 @@ func (vector Vector) String() string { // MarshalBinary implements encoding.BinaryMarshaler -func (vector *Vector) MarshalBinary() (data []byte, err error) { +func (vector *VectorE4) MarshalBinary() (data []byte, err error) { var buf bytes.Buffer if _, err = vector.WriteTo(&buf); err != nil { @@ -416,7 +421,7 @@ func (vector *Vector) MarshalBinary() (data []byte, err error) { } // UnmarshalBinary implements encoding.BinaryUnmarshaler -func (vector *Vector) UnmarshalBinary(data []byte) error { +func (vector *VectorE4) UnmarshalBinary(data []byte) error { r := bytes.NewReader(data) _, err := vector.ReadFrom(r) return err @@ -424,7 +429,7 @@ func (vector *Vector) UnmarshalBinary(data []byte) error { // WriteTo implements io.WriterTo and writes a vector of big endian encoded Element. // Length of the vector is encoded as a uint32 on the first 4 bytes. -func (vector *Vector) WriteTo(w io.Writer) (int64, error) { +func (vector *VectorE4) WriteTo(w io.Writer) (int64, error) { // encode slice length if err := binary.Write(w, binary.BigEndian, uint32(len(*vector))); err != nil { @@ -452,7 +457,7 @@ func (vector *Vector) WriteTo(w io.Writer) (int64, error) { return n, nil } -// AsyncReadFrom implements an asynchronous version of [Vector.ReadFrom]. It +// AsyncReadFrom implements an asynchronous version of [VectorE4.ReadFrom]. It // reads the reader r in full and then performs the validation and conversion to // Montgomery form separately in a goroutine. Any error encountered during // reading is returned directly, while errors encountered during @@ -476,7 +481,7 @@ func (vector *Vector) WriteTo(w io.Writer) (int64, error) { // - first 4 bytes: length of the vector as a big-endian uint32 // - for each element of the vector, `4 * fr.Bytes` bytes representing the // element in big-endian encoding. -func (vector *Vector) AsyncReadFrom(r io.Reader) (int64, error, chan error) { // nolint ST1008 +func (vector *VectorE4) AsyncReadFrom(r io.Reader) (int64, error, chan error) { // nolint ST1008 chErr := make(chan error, 1) var bufSizeSlice [4]byte if read, err := io.ReadFull(r, bufSizeSlice[:]); err != nil { @@ -505,7 +510,7 @@ func (vector *Vector) AsyncReadFrom(r io.Reader) (int64, error, chan error) { // *vector = (*vector)[:0] if headerSliceLen == 0 { if *vector == nil { - *vector = Vector{} + *vector = VectorE4{} } close(chErr) return totalRead, nil, chErr @@ -513,7 +518,7 @@ func (vector *Vector) AsyncReadFrom(r io.Reader) (int64, error, chan error) { // for i := uint64(0); i < headerSliceLen; i += maxAllocateSliceLength { if len(*vector) <= int(i) { - *vector = append(*vector, make(Vector, int(min(headerSliceLen-i, maxAllocateSliceLength)))...) + *vector = append(*vector, make(VectorE4, int(min(headerSliceLen-i, maxAllocateSliceLength)))...) } bSlice := unsafe.Slice((*byte)(unsafe.Pointer(&(*vector)[i])), int(min(headerSliceLen-i, maxAllocateSliceLength))*BytesE4) read, err := io.ReadFull(r, bSlice) @@ -586,7 +591,7 @@ func (vector *Vector) AsyncReadFrom(r io.Reader) (int64, error, chan error) { // // - for each element of the vector, `4 * fr.Bytes` bytes representing the element in big-endian encoding. // // The method implements [io.ReaderFrom] interface. -func (vector *Vector) ReadFrom(r io.Reader) (int64, error) { +func (vector *VectorE4) ReadFrom(r io.Reader) (int64, error) { // call the async version and wait for the channel to be closed n, err, chErr := vector.AsyncReadFrom(r) @@ -600,7 +605,7 @@ func (vector *Vector) ReadFrom(r io.Reader) (int64, error) { // i.e. vector[i] = vector[0] * vector[1] * ... * vector[i] // If nbTasks > 1, it uses nbTasks goroutines to compute the prefix product in parallel. // If nbTasks is not provided, it uses the number of CPU cores. -func (vector Vector) PrefixProduct(nbTasks ...int) { +func (vector VectorE4) PrefixProduct(nbTasks ...int) { N := len(vector) if N < 2 { return @@ -660,7 +665,7 @@ func (vector Vector) PrefixProduct(nbTasks ...int) { } -func (vector Vector) prefixProductGeneric() { +func (vector VectorE4) prefixProductGeneric() { for i := 1; i < len(vector); i++ { vector[i].Mul(&vector[i], &vector[i-1]) } @@ -668,28 +673,28 @@ func (vector Vector) prefixProductGeneric() { -func vectorAddGeneric(res, a, b Vector) { +func vectorAddGeneric(res, a, b VectorE4) { for i := range len(res) { res[i].Add(&a[i], &b[i]) } } -func vectorSubGeneric(res, a, b Vector) { +func vectorSubGeneric(res, a, b VectorE4) { for i := range len(res) { res[i].Sub(&a[i], &b[i]) } } -func vectorMulGeneric(res, a, b Vector) { +func vectorMulGeneric(res, a, b VectorE4) { for i := range len(res) { res[i].Mul(&a[i], &b[i]) } } -func vectorScalarMulGeneric(res, a Vector, b *E4) { +func vectorScalarMulGeneric(res, a VectorE4, b *E4) { for i := range len(res) { res[i].Mul(&a[i], b) } } -func vectorInnerProductGeneric(a, b Vector) E4 { +func vectorInnerProductGeneric(a, b VectorE4) E4 { var res, tmp E4 for i := range len(a) { tmp.Mul(&a[i], &b[i]) @@ -698,7 +703,7 @@ func vectorInnerProductGeneric(a, b Vector) E4 { return res } -func vectorInnerProductByElementGeneric(a Vector, b fr.Vector) E4 { +func vectorInnerProductByElementGeneric(a VectorE4, b fr.Vector) E4 { var res, tmp E4 for i := range len(a) { tmp.MulByElement(&a[i], &b[i]) @@ -707,7 +712,7 @@ func vectorInnerProductByElementGeneric(a Vector, b fr.Vector) E4 { return res } -func vectorSumGeneric(v Vector) E4 { +func vectorSumGeneric(v VectorE4) E4 { var sum E4 for i := range len(v) { sum.Add(&sum, &v[i]) @@ -715,7 +720,7 @@ func vectorSumGeneric(v Vector) E4 { return sum } -func vectorMulAccByElementGeneric(v Vector, scale []fr.Element, alpha *E4) { +func vectorMulAccByElementGeneric(v VectorE4, scale []fr.Element, alpha *E4) { var tmp E4 for i := range len(v) { tmp.MulByElement(alpha, &scale[i]) @@ -723,13 +728,13 @@ func vectorMulAccByElementGeneric(v Vector, scale []fr.Element, alpha *E4) { } } -func vectorMulByElementGeneric(res, a Vector, b fr.Vector) { +func vectorMulByElementGeneric(res, a VectorE4, b fr.Vector) { for i := range len(res) { res[i].MulByElement(&a[i], &b[i]) } } -func vectorButterflyGeneric(a, b Vector) { +func vectorButterflyGeneric(a, b VectorE4) { for i := range len(a) { Butterfly(&a[i], &b[i]) } diff --git a/internal/generator/field/template/fft/fftext.go.tmpl b/internal/generator/field/template/fft/fftext.go.tmpl index d34f1cc31a..1308aed845 100644 --- a/internal/generator/field/template/fft/fftext.go.tmpl +++ b/internal/generator/field/template/fft/fftext.go.tmpl @@ -54,7 +54,7 @@ func (domain *Domain) FFTExt(a []fext.{{.ExtType}}, decimation Decimation, opts } } parallel.ExecuteAligned(len(a), {{.ExtAlign}}, func(start, end int) { - va := fext.Vector(a[start:end]) + va := fext.Vector{{.ExtType}}(a[start:end]) va.MulByElement(va, cosetTable[start:end]) }, opt.nbTasks) } @@ -124,7 +124,7 @@ func (domain *Domain) FFTInverseExt(a []fext.{{.ExtType}}, decimation Decimation // scale by CardinalityInv if !opt.coset { parallel.ExecuteAligned(len(a), {{.ExtAlign}}, func(start, end int) { - va := fext.Vector(a[start:end]) + va := fext.Vector{{.ExtType}}(a[start:end]) va.ScalarMulByElement(va, &domain.CardinalityInv) }, opt.nbTasks) return @@ -154,7 +154,7 @@ func (domain *Domain) FFTInverseExt(a []fext.{{.ExtType}}, decimation Decimation } } parallel.ExecuteAligned(len(a), {{.ExtAlign}}, func(start, end int) { - va := fext.Vector(a[start:end]) + va := fext.Vector{{.ExtType}}(a[start:end]) va.MulByElement(va, cosetTableInv[start:end]) va.ScalarMulByElement(va, &domain.CardinalityInv) }, opt.nbTasks) @@ -225,8 +225,8 @@ func difFFTExt(a []fext.{{.ExtType}}, w {{ .FF }}.Element, twiddles [][]{{ .FF } } func innerDIFWithTwiddlesExt(a []fext.{{.ExtType}}, twiddles []{{ .FF }}.Element, start, end, m int) { - va0 := fext.Vector(a[start:end]) - va1 := fext.Vector(a[start+m : end+m]) + va0 := fext.Vector{{.ExtType}}(a[start:end]) + va1 := fext.Vector{{.ExtType}}(a[start+m : end+m]) va0.Butterfly(va1) va1.MulByElement(va1, twiddles[start:end]) } @@ -308,8 +308,8 @@ func ditFFTExt(a []fext.{{.ExtType}}, w {{ .FF }}.Element, twiddles [][]{{ .FF } func innerDITWithTwiddlesExt(a []fext.{{.ExtType}}, twiddles []{{ .FF }}.Element, start, end, m int) { - va0 := fext.Vector(a[start:end]) - va1 := fext.Vector(a[start+m : end+m]) + va0 := fext.Vector{{.ExtType}}(a[start:end]) + va1 := fext.Vector{{.ExtType}}(a[start+m : end+m]) va1.MulByElement(va1, twiddles[start:end]) va0.Butterfly(va1) } @@ -352,7 +352,7 @@ func kerDIFNP_{{.sizeKernel}}Ext(a []fext.{{.ExtType}}, twiddles [][]{{ .FF }}.E innerDIFWithTwiddlesExt(a[:{{$n}}], twiddles[stage + {{$step}}], 0, {{$m}}, {{$m}}) {{- else}} {{- if eq $m 1}} - va := fext.Vector(a[:{{$bound}}]) + va := fext.Vector{{$.ExtType}}(a[:{{$bound}}]) va.ButterflyPair() {{- else}} for offset := 0; offset < {{$bound}}; offset += {{$n}} { @@ -380,7 +380,7 @@ func kerDITNP_{{.sizeKernel}}Ext(a []fext.{{.ExtType}}, twiddles [][]{{ .FF }}.E innerDITWithTwiddlesExt(a[:{{$n}}], twiddles[stage + {{$step}}], 0, {{$m}}, {{$m}}) {{- else}} {{- if eq $m 1}} - va := fext.Vector(a[:{{$bound}}]) + va := fext.Vector{{$.ExtType}}(a[:{{$bound}}]) va.ButterflyPair() {{- else if eq $m 2}} for offset := 0; offset < {{$bound}}; offset += {{$n}} { From 968014f6f7ead4584af4a3144e1ca7bda3b4a69f Mon Sep 17 00:00:00 2001 From: Arya Tabaie Date: Mon, 5 Oct 2026 15:17:45 -0500 Subject: [PATCH 09/18] refactor: SetBytesCanonical instead of SetBytes Signed-off-by: Arya Tabaie --- field/babybear/extensions/e2.go | 21 ++++++---- field/babybear/extensions/e2_test.go | 27 ++++++++++--- field/babybear/extensions/e4.go | 29 +++++++++----- field/babybear/extensions/e4_test.go | 27 ++++++++++--- field/babybear/extensions/e6.go | 39 +++++++++++++------ field/babybear/extensions/e6_test.go | 27 ++++++++++--- field/goldilocks/extensions/e2.go | 21 ++++++---- field/goldilocks/extensions/e2_test.go | 27 ++++++++++--- field/koalabear/extensions/e2.go | 21 ++++++---- field/koalabear/extensions/e2_test.go | 27 ++++++++++--- field/koalabear/extensions/e4.go | 29 +++++++++----- field/koalabear/extensions/e4_test.go | 27 ++++++++++--- field/koalabear/extensions/e6.go | 39 +++++++++++++------ field/koalabear/extensions/e6_test.go | 27 ++++++++++--- field/mamabear/extensions/e3.go | 25 ++++++++---- field/mamabear/extensions/e3_test.go | 27 ++++++++++--- .../field/template/extensions/e2.go.tmpl | 21 ++++++---- .../field/template/extensions/e2_test.go.tmpl | 27 ++++++++++--- .../field/template/extensions/e3.go.tmpl | 25 ++++++++---- .../field/template/extensions/e3_test.go.tmpl | 27 ++++++++++--- .../field/template/extensions/e4.go.tmpl | 29 +++++++++----- .../field/template/extensions/e4_test.go.tmpl | 27 ++++++++++--- .../field/template/extensions/e6.go.tmpl | 39 +++++++++++++------ .../field/template/extensions/e6_test.go.tmpl | 27 ++++++++++--- 24 files changed, 483 insertions(+), 179 deletions(-) diff --git a/field/babybear/extensions/e2.go b/field/babybear/extensions/e2.go index 8672c2de42..56371d4912 100644 --- a/field/babybear/extensions/e2.go +++ b/field/babybear/extensions/e2.go @@ -114,15 +114,22 @@ func (z *E2) Marshal() []byte { return res } -// SetBytes sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytes, and returns z. -// It returns an error if len(b) != BytesE2. -func (z *E2) SetBytes(b []byte) (*E2, error) { +// SetBytesCanonical sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytesCanonical. +// It returns an error if len(b) != BytesE2 or if any coefficient is not the canonical encoding of a field element; +// in that case z is left unchanged. +func (z *E2) SetBytesCanonical(b []byte) error { if len(b) != BytesE2 { - return nil, fmt.Errorf("E2.SetBytes: got %d bytes, expected %d", len(b), BytesE2) + return fmt.Errorf("E2.SetBytesCanonical: got %d bytes, expected %d", len(b), BytesE2) } - z.A0.SetBytes(b[0*fr.Bytes : 1*fr.Bytes]) - z.A1.SetBytes(b[1*fr.Bytes : 2*fr.Bytes]) - return z, nil + var r E2 + if err := r.A0.SetBytesCanonical(b[0*fr.Bytes : 1*fr.Bytes]); err != nil { + return err + } + if err := r.A1.SetBytesCanonical(b[1*fr.Bytes : 2*fr.Bytes]); err != nil { + return err + } + *z = r + return nil } // SetRandom sets a0 and a1 to random values diff --git a/field/babybear/extensions/e2_test.go b/field/babybear/extensions/e2_test.go index da9cb653ef..37e3d7d4c1 100644 --- a/field/babybear/extensions/e2_test.go +++ b/field/babybear/extensions/e2_test.go @@ -557,7 +557,7 @@ func genE2() gopter.Gen { }) } -func TestE2MarshalSetBytesRoundTrip(t *testing.T) { +func TestE2MarshalSetBytesCanonicalRoundTrip(t *testing.T) { for range 100 { var x E2 x.MustSetRandom() @@ -575,17 +575,16 @@ func TestE2MarshalSetBytesRoundTrip(t *testing.T) { } var y E2 - _, err := y.SetBytes(b) - require.NoError(t, err) + require.NoError(t, y.SetBytesCanonical(b)) if !y.Equal(&x) { - t.Fatal("SetBytes(Marshal(x)) != x") + t.Fatal("SetBytesCanonical(Marshal(x)) != x") } } for _, n := range []int{0, BytesE2 - 1, BytesE2 + 1} { var z E2 - if _, err := z.SetBytes(make([]byte, n)); err == nil { - t.Fatalf("SetBytes did not fail on %d bytes", n) + if err := z.SetBytesCanonical(make([]byte, n)); err == nil { + t.Fatalf("SetBytesCanonical did not fail on %d bytes", n) } } } @@ -617,3 +616,19 @@ func TestE2SetInt64SetUint64(t *testing.T) { require.Truef(t, z.Equal(&zb), "SetUint64(%d) != SetBigInt(%d)", v, v) } } + +func TestE2SetBytesCanonicalRejectsNonCanonical(t *testing.T) { + var x E2 + x.MustSetRandom() + good := x.Marshal() + + q := fr.Modulus() + for i := range 2 { + b := append([]byte(nil), good...) + q.FillBytes(b[i*fr.Bytes : (i+1)*fr.Bytes]) + + y := x + require.Error(t, y.SetBytesCanonical(b), "coefficient %d equal to the modulus must be rejected", i) + require.True(t, y.Equal(&x), "E2 must be unchanged on error") + } +} diff --git a/field/babybear/extensions/e4.go b/field/babybear/extensions/e4.go index 811a714985..879dc452f6 100644 --- a/field/babybear/extensions/e4.go +++ b/field/babybear/extensions/e4.go @@ -127,17 +127,28 @@ func (z *E4) Marshal() []byte { return res } -// SetBytes sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytes, and returns z. -// It returns an error if len(b) != BytesE4. -func (z *E4) SetBytes(b []byte) (*E4, error) { +// SetBytesCanonical sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytesCanonical. +// It returns an error if len(b) != BytesE4 or if any coefficient is not the canonical encoding of a field element; +// in that case z is left unchanged. +func (z *E4) SetBytesCanonical(b []byte) error { if len(b) != BytesE4 { - return nil, fmt.Errorf("E4.SetBytes: got %d bytes, expected %d", len(b), BytesE4) + return fmt.Errorf("E4.SetBytesCanonical: got %d bytes, expected %d", len(b), BytesE4) } - z.B0.A0.SetBytes(b[0*fr.Bytes : 1*fr.Bytes]) - z.B0.A1.SetBytes(b[1*fr.Bytes : 2*fr.Bytes]) - z.B1.A0.SetBytes(b[2*fr.Bytes : 3*fr.Bytes]) - z.B1.A1.SetBytes(b[3*fr.Bytes : 4*fr.Bytes]) - return z, nil + var r E4 + if err := r.B0.A0.SetBytesCanonical(b[0*fr.Bytes : 1*fr.Bytes]); err != nil { + return err + } + if err := r.B0.A1.SetBytesCanonical(b[1*fr.Bytes : 2*fr.Bytes]); err != nil { + return err + } + if err := r.B1.A0.SetBytesCanonical(b[2*fr.Bytes : 3*fr.Bytes]); err != nil { + return err + } + if err := r.B1.A1.SetBytesCanonical(b[3*fr.Bytes : 4*fr.Bytes]); err != nil { + return err + } + *z = r + return nil } // Lift sets the B0.A0 component of z to v diff --git a/field/babybear/extensions/e4_test.go b/field/babybear/extensions/e4_test.go index 687f317e48..2e514ab2dc 100644 --- a/field/babybear/extensions/e4_test.go +++ b/field/babybear/extensions/e4_test.go @@ -1127,7 +1127,7 @@ func genFrVector(size int) gopter.Gen { } } -func TestE4MarshalSetBytesRoundTrip(t *testing.T) { +func TestE4MarshalSetBytesCanonicalRoundTrip(t *testing.T) { for range 100 { var x E4 x.MustSetRandom() @@ -1145,17 +1145,16 @@ func TestE4MarshalSetBytesRoundTrip(t *testing.T) { } var y E4 - _, err := y.SetBytes(b) - require.NoError(t, err) + require.NoError(t, y.SetBytesCanonical(b)) if !y.Equal(&x) { - t.Fatal("SetBytes(Marshal(x)) != x") + t.Fatal("SetBytesCanonical(Marshal(x)) != x") } } for _, n := range []int{0, BytesE4 - 1, BytesE4 + 1} { var z E4 - if _, err := z.SetBytes(make([]byte, n)); err == nil { - t.Fatalf("SetBytes did not fail on %d bytes", n) + if err := z.SetBytesCanonical(make([]byte, n)); err == nil { + t.Fatalf("SetBytesCanonical did not fail on %d bytes", n) } } } @@ -1191,3 +1190,19 @@ func TestE4SetInt64SetUint64(t *testing.T) { require.Truef(t, z.Equal(&zb), "SetUint64(%d) != SetBigInt(%d)", v, v) } } + +func TestE4SetBytesCanonicalRejectsNonCanonical(t *testing.T) { + var x E4 + x.MustSetRandom() + good := x.Marshal() + + q := fr.Modulus() + for i := range 4 { + b := append([]byte(nil), good...) + q.FillBytes(b[i*fr.Bytes : (i+1)*fr.Bytes]) + + y := x + require.Error(t, y.SetBytesCanonical(b), "coefficient %d equal to the modulus must be rejected", i) + require.True(t, y.Equal(&x), "E4 must be unchanged on error") + } +} diff --git a/field/babybear/extensions/e6.go b/field/babybear/extensions/e6.go index 38e5f41c75..f9569184a4 100644 --- a/field/babybear/extensions/e6.go +++ b/field/babybear/extensions/e6.go @@ -119,19 +119,34 @@ func (z *E6) Marshal() []byte { return res } -// SetBytes sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytes, and returns z. -// It returns an error if len(b) != BytesE6. -func (z *E6) SetBytes(b []byte) (*E6, error) { +// SetBytesCanonical sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytesCanonical. +// It returns an error if len(b) != BytesE6 or if any coefficient is not the canonical encoding of a field element; +// in that case z is left unchanged. +func (z *E6) SetBytesCanonical(b []byte) error { if len(b) != BytesE6 { - return nil, fmt.Errorf("E6.SetBytes: got %d bytes, expected %d", len(b), BytesE6) - } - z.B0.A0.SetBytes(b[0*fr.Bytes : 1*fr.Bytes]) - z.B0.A1.SetBytes(b[1*fr.Bytes : 2*fr.Bytes]) - z.B1.A0.SetBytes(b[2*fr.Bytes : 3*fr.Bytes]) - z.B1.A1.SetBytes(b[3*fr.Bytes : 4*fr.Bytes]) - z.B2.A0.SetBytes(b[4*fr.Bytes : 5*fr.Bytes]) - z.B2.A1.SetBytes(b[5*fr.Bytes : 6*fr.Bytes]) - return z, nil + return fmt.Errorf("E6.SetBytesCanonical: got %d bytes, expected %d", len(b), BytesE6) + } + var r E6 + if err := r.B0.A0.SetBytesCanonical(b[0*fr.Bytes : 1*fr.Bytes]); err != nil { + return err + } + if err := r.B0.A1.SetBytesCanonical(b[1*fr.Bytes : 2*fr.Bytes]); err != nil { + return err + } + if err := r.B1.A0.SetBytesCanonical(b[2*fr.Bytes : 3*fr.Bytes]); err != nil { + return err + } + if err := r.B1.A1.SetBytesCanonical(b[3*fr.Bytes : 4*fr.Bytes]); err != nil { + return err + } + if err := r.B2.A0.SetBytesCanonical(b[4*fr.Bytes : 5*fr.Bytes]); err != nil { + return err + } + if err := r.B2.A1.SetBytesCanonical(b[5*fr.Bytes : 6*fr.Bytes]); err != nil { + return err + } + *z = r + return nil } // MulByElement multiplies an element in E6 by an element in fr. diff --git a/field/babybear/extensions/e6_test.go b/field/babybear/extensions/e6_test.go index 2bf902dd5c..91a0a10d0b 100644 --- a/field/babybear/extensions/e6_test.go +++ b/field/babybear/extensions/e6_test.go @@ -362,7 +362,7 @@ func genE6() gopter.Gen { }) } -func TestE6MarshalSetBytesRoundTrip(t *testing.T) { +func TestE6MarshalSetBytesCanonicalRoundTrip(t *testing.T) { for range 100 { var x E6 x.MustSetRandom() @@ -380,17 +380,32 @@ func TestE6MarshalSetBytesRoundTrip(t *testing.T) { } var y E6 - _, err := y.SetBytes(b) - require.NoError(t, err) + require.NoError(t, y.SetBytesCanonical(b)) if !y.Equal(&x) { - t.Fatal("SetBytes(Marshal(x)) != x") + t.Fatal("SetBytesCanonical(Marshal(x)) != x") } } for _, n := range []int{0, BytesE6 - 1, BytesE6 + 1} { var z E6 - if _, err := z.SetBytes(make([]byte, n)); err == nil { - t.Fatalf("SetBytes did not fail on %d bytes", n) + if err := z.SetBytesCanonical(make([]byte, n)); err == nil { + t.Fatalf("SetBytesCanonical did not fail on %d bytes", n) } } } + +func TestE6SetBytesCanonicalRejectsNonCanonical(t *testing.T) { + var x E6 + x.MustSetRandom() + good := x.Marshal() + + q := fr.Modulus() + for i := range 6 { + b := append([]byte(nil), good...) + q.FillBytes(b[i*fr.Bytes : (i+1)*fr.Bytes]) + + y := x + require.Error(t, y.SetBytesCanonical(b), "coefficient %d equal to the modulus must be rejected", i) + require.True(t, y.Equal(&x), "E6 must be unchanged on error") + } +} diff --git a/field/goldilocks/extensions/e2.go b/field/goldilocks/extensions/e2.go index 5ea33e19fa..35345a29c6 100644 --- a/field/goldilocks/extensions/e2.go +++ b/field/goldilocks/extensions/e2.go @@ -114,15 +114,22 @@ func (z *E2) Marshal() []byte { return res } -// SetBytes sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytes, and returns z. -// It returns an error if len(b) != BytesE2. -func (z *E2) SetBytes(b []byte) (*E2, error) { +// SetBytesCanonical sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytesCanonical. +// It returns an error if len(b) != BytesE2 or if any coefficient is not the canonical encoding of a field element; +// in that case z is left unchanged. +func (z *E2) SetBytesCanonical(b []byte) error { if len(b) != BytesE2 { - return nil, fmt.Errorf("E2.SetBytes: got %d bytes, expected %d", len(b), BytesE2) + return fmt.Errorf("E2.SetBytesCanonical: got %d bytes, expected %d", len(b), BytesE2) } - z.A0.SetBytes(b[0*fr.Bytes : 1*fr.Bytes]) - z.A1.SetBytes(b[1*fr.Bytes : 2*fr.Bytes]) - return z, nil + var r E2 + if err := r.A0.SetBytesCanonical(b[0*fr.Bytes : 1*fr.Bytes]); err != nil { + return err + } + if err := r.A1.SetBytesCanonical(b[1*fr.Bytes : 2*fr.Bytes]); err != nil { + return err + } + *z = r + return nil } // SetRandom sets a0 and a1 to random values diff --git a/field/goldilocks/extensions/e2_test.go b/field/goldilocks/extensions/e2_test.go index 7c10621772..cda8d7adb5 100644 --- a/field/goldilocks/extensions/e2_test.go +++ b/field/goldilocks/extensions/e2_test.go @@ -540,7 +540,7 @@ func genE2() gopter.Gen { }) } -func TestE2MarshalSetBytesRoundTrip(t *testing.T) { +func TestE2MarshalSetBytesCanonicalRoundTrip(t *testing.T) { for range 100 { var x E2 x.MustSetRandom() @@ -558,17 +558,16 @@ func TestE2MarshalSetBytesRoundTrip(t *testing.T) { } var y E2 - _, err := y.SetBytes(b) - require.NoError(t, err) + require.NoError(t, y.SetBytesCanonical(b)) if !y.Equal(&x) { - t.Fatal("SetBytes(Marshal(x)) != x") + t.Fatal("SetBytesCanonical(Marshal(x)) != x") } } for _, n := range []int{0, BytesE2 - 1, BytesE2 + 1} { var z E2 - if _, err := z.SetBytes(make([]byte, n)); err == nil { - t.Fatalf("SetBytes did not fail on %d bytes", n) + if err := z.SetBytesCanonical(make([]byte, n)); err == nil { + t.Fatalf("SetBytesCanonical did not fail on %d bytes", n) } } } @@ -600,3 +599,19 @@ func TestE2SetInt64SetUint64(t *testing.T) { require.Truef(t, z.Equal(&zb), "SetUint64(%d) != SetBigInt(%d)", v, v) } } + +func TestE2SetBytesCanonicalRejectsNonCanonical(t *testing.T) { + var x E2 + x.MustSetRandom() + good := x.Marshal() + + q := fr.Modulus() + for i := range 2 { + b := append([]byte(nil), good...) + q.FillBytes(b[i*fr.Bytes : (i+1)*fr.Bytes]) + + y := x + require.Error(t, y.SetBytesCanonical(b), "coefficient %d equal to the modulus must be rejected", i) + require.True(t, y.Equal(&x), "E2 must be unchanged on error") + } +} diff --git a/field/koalabear/extensions/e2.go b/field/koalabear/extensions/e2.go index fa89720ee7..db49bafcf7 100644 --- a/field/koalabear/extensions/e2.go +++ b/field/koalabear/extensions/e2.go @@ -114,15 +114,22 @@ func (z *E2) Marshal() []byte { return res } -// SetBytes sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytes, and returns z. -// It returns an error if len(b) != BytesE2. -func (z *E2) SetBytes(b []byte) (*E2, error) { +// SetBytesCanonical sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytesCanonical. +// It returns an error if len(b) != BytesE2 or if any coefficient is not the canonical encoding of a field element; +// in that case z is left unchanged. +func (z *E2) SetBytesCanonical(b []byte) error { if len(b) != BytesE2 { - return nil, fmt.Errorf("E2.SetBytes: got %d bytes, expected %d", len(b), BytesE2) + return fmt.Errorf("E2.SetBytesCanonical: got %d bytes, expected %d", len(b), BytesE2) } - z.A0.SetBytes(b[0*fr.Bytes : 1*fr.Bytes]) - z.A1.SetBytes(b[1*fr.Bytes : 2*fr.Bytes]) - return z, nil + var r E2 + if err := r.A0.SetBytesCanonical(b[0*fr.Bytes : 1*fr.Bytes]); err != nil { + return err + } + if err := r.A1.SetBytesCanonical(b[1*fr.Bytes : 2*fr.Bytes]); err != nil { + return err + } + *z = r + return nil } // SetRandom sets a0 and a1 to random values diff --git a/field/koalabear/extensions/e2_test.go b/field/koalabear/extensions/e2_test.go index 8ee2dd1f18..df462eab32 100644 --- a/field/koalabear/extensions/e2_test.go +++ b/field/koalabear/extensions/e2_test.go @@ -557,7 +557,7 @@ func genE2() gopter.Gen { }) } -func TestE2MarshalSetBytesRoundTrip(t *testing.T) { +func TestE2MarshalSetBytesCanonicalRoundTrip(t *testing.T) { for range 100 { var x E2 x.MustSetRandom() @@ -575,17 +575,16 @@ func TestE2MarshalSetBytesRoundTrip(t *testing.T) { } var y E2 - _, err := y.SetBytes(b) - require.NoError(t, err) + require.NoError(t, y.SetBytesCanonical(b)) if !y.Equal(&x) { - t.Fatal("SetBytes(Marshal(x)) != x") + t.Fatal("SetBytesCanonical(Marshal(x)) != x") } } for _, n := range []int{0, BytesE2 - 1, BytesE2 + 1} { var z E2 - if _, err := z.SetBytes(make([]byte, n)); err == nil { - t.Fatalf("SetBytes did not fail on %d bytes", n) + if err := z.SetBytesCanonical(make([]byte, n)); err == nil { + t.Fatalf("SetBytesCanonical did not fail on %d bytes", n) } } } @@ -617,3 +616,19 @@ func TestE2SetInt64SetUint64(t *testing.T) { require.Truef(t, z.Equal(&zb), "SetUint64(%d) != SetBigInt(%d)", v, v) } } + +func TestE2SetBytesCanonicalRejectsNonCanonical(t *testing.T) { + var x E2 + x.MustSetRandom() + good := x.Marshal() + + q := fr.Modulus() + for i := range 2 { + b := append([]byte(nil), good...) + q.FillBytes(b[i*fr.Bytes : (i+1)*fr.Bytes]) + + y := x + require.Error(t, y.SetBytesCanonical(b), "coefficient %d equal to the modulus must be rejected", i) + require.True(t, y.Equal(&x), "E2 must be unchanged on error") + } +} diff --git a/field/koalabear/extensions/e4.go b/field/koalabear/extensions/e4.go index c164108368..31af799017 100644 --- a/field/koalabear/extensions/e4.go +++ b/field/koalabear/extensions/e4.go @@ -127,17 +127,28 @@ func (z *E4) Marshal() []byte { return res } -// SetBytes sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytes, and returns z. -// It returns an error if len(b) != BytesE4. -func (z *E4) SetBytes(b []byte) (*E4, error) { +// SetBytesCanonical sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytesCanonical. +// It returns an error if len(b) != BytesE4 or if any coefficient is not the canonical encoding of a field element; +// in that case z is left unchanged. +func (z *E4) SetBytesCanonical(b []byte) error { if len(b) != BytesE4 { - return nil, fmt.Errorf("E4.SetBytes: got %d bytes, expected %d", len(b), BytesE4) + return fmt.Errorf("E4.SetBytesCanonical: got %d bytes, expected %d", len(b), BytesE4) } - z.B0.A0.SetBytes(b[0*fr.Bytes : 1*fr.Bytes]) - z.B0.A1.SetBytes(b[1*fr.Bytes : 2*fr.Bytes]) - z.B1.A0.SetBytes(b[2*fr.Bytes : 3*fr.Bytes]) - z.B1.A1.SetBytes(b[3*fr.Bytes : 4*fr.Bytes]) - return z, nil + var r E4 + if err := r.B0.A0.SetBytesCanonical(b[0*fr.Bytes : 1*fr.Bytes]); err != nil { + return err + } + if err := r.B0.A1.SetBytesCanonical(b[1*fr.Bytes : 2*fr.Bytes]); err != nil { + return err + } + if err := r.B1.A0.SetBytesCanonical(b[2*fr.Bytes : 3*fr.Bytes]); err != nil { + return err + } + if err := r.B1.A1.SetBytesCanonical(b[3*fr.Bytes : 4*fr.Bytes]); err != nil { + return err + } + *z = r + return nil } // Lift sets the B0.A0 component of z to v diff --git a/field/koalabear/extensions/e4_test.go b/field/koalabear/extensions/e4_test.go index d15b36792d..c55631d70f 100644 --- a/field/koalabear/extensions/e4_test.go +++ b/field/koalabear/extensions/e4_test.go @@ -1127,7 +1127,7 @@ func genFrVector(size int) gopter.Gen { } } -func TestE4MarshalSetBytesRoundTrip(t *testing.T) { +func TestE4MarshalSetBytesCanonicalRoundTrip(t *testing.T) { for range 100 { var x E4 x.MustSetRandom() @@ -1145,17 +1145,16 @@ func TestE4MarshalSetBytesRoundTrip(t *testing.T) { } var y E4 - _, err := y.SetBytes(b) - require.NoError(t, err) + require.NoError(t, y.SetBytesCanonical(b)) if !y.Equal(&x) { - t.Fatal("SetBytes(Marshal(x)) != x") + t.Fatal("SetBytesCanonical(Marshal(x)) != x") } } for _, n := range []int{0, BytesE4 - 1, BytesE4 + 1} { var z E4 - if _, err := z.SetBytes(make([]byte, n)); err == nil { - t.Fatalf("SetBytes did not fail on %d bytes", n) + if err := z.SetBytesCanonical(make([]byte, n)); err == nil { + t.Fatalf("SetBytesCanonical did not fail on %d bytes", n) } } } @@ -1191,3 +1190,19 @@ func TestE4SetInt64SetUint64(t *testing.T) { require.Truef(t, z.Equal(&zb), "SetUint64(%d) != SetBigInt(%d)", v, v) } } + +func TestE4SetBytesCanonicalRejectsNonCanonical(t *testing.T) { + var x E4 + x.MustSetRandom() + good := x.Marshal() + + q := fr.Modulus() + for i := range 4 { + b := append([]byte(nil), good...) + q.FillBytes(b[i*fr.Bytes : (i+1)*fr.Bytes]) + + y := x + require.Error(t, y.SetBytesCanonical(b), "coefficient %d equal to the modulus must be rejected", i) + require.True(t, y.Equal(&x), "E4 must be unchanged on error") + } +} diff --git a/field/koalabear/extensions/e6.go b/field/koalabear/extensions/e6.go index 6a2c755d4a..c7df7659d0 100644 --- a/field/koalabear/extensions/e6.go +++ b/field/koalabear/extensions/e6.go @@ -120,19 +120,34 @@ func (z *E6) Marshal() []byte { return res } -// SetBytes sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytes, and returns z. -// It returns an error if len(b) != BytesE6. -func (z *E6) SetBytes(b []byte) (*E6, error) { +// SetBytesCanonical sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytesCanonical. +// It returns an error if len(b) != BytesE6 or if any coefficient is not the canonical encoding of a field element; +// in that case z is left unchanged. +func (z *E6) SetBytesCanonical(b []byte) error { if len(b) != BytesE6 { - return nil, fmt.Errorf("E6.SetBytes: got %d bytes, expected %d", len(b), BytesE6) - } - z.B0.A0.SetBytes(b[0*fr.Bytes : 1*fr.Bytes]) - z.B0.A1.SetBytes(b[1*fr.Bytes : 2*fr.Bytes]) - z.B1.A0.SetBytes(b[2*fr.Bytes : 3*fr.Bytes]) - z.B1.A1.SetBytes(b[3*fr.Bytes : 4*fr.Bytes]) - z.B2.A0.SetBytes(b[4*fr.Bytes : 5*fr.Bytes]) - z.B2.A1.SetBytes(b[5*fr.Bytes : 6*fr.Bytes]) - return z, nil + return fmt.Errorf("E6.SetBytesCanonical: got %d bytes, expected %d", len(b), BytesE6) + } + var r E6 + if err := r.B0.A0.SetBytesCanonical(b[0*fr.Bytes : 1*fr.Bytes]); err != nil { + return err + } + if err := r.B0.A1.SetBytesCanonical(b[1*fr.Bytes : 2*fr.Bytes]); err != nil { + return err + } + if err := r.B1.A0.SetBytesCanonical(b[2*fr.Bytes : 3*fr.Bytes]); err != nil { + return err + } + if err := r.B1.A1.SetBytesCanonical(b[3*fr.Bytes : 4*fr.Bytes]); err != nil { + return err + } + if err := r.B2.A0.SetBytesCanonical(b[4*fr.Bytes : 5*fr.Bytes]); err != nil { + return err + } + if err := r.B2.A1.SetBytesCanonical(b[5*fr.Bytes : 6*fr.Bytes]); err != nil { + return err + } + *z = r + return nil } // MulByElement multiplies an element in E6 by an element in fr. diff --git a/field/koalabear/extensions/e6_test.go b/field/koalabear/extensions/e6_test.go index 27129696df..02187bcabc 100644 --- a/field/koalabear/extensions/e6_test.go +++ b/field/koalabear/extensions/e6_test.go @@ -362,7 +362,7 @@ func genE6() gopter.Gen { }) } -func TestE6MarshalSetBytesRoundTrip(t *testing.T) { +func TestE6MarshalSetBytesCanonicalRoundTrip(t *testing.T) { for range 100 { var x E6 x.MustSetRandom() @@ -380,17 +380,32 @@ func TestE6MarshalSetBytesRoundTrip(t *testing.T) { } var y E6 - _, err := y.SetBytes(b) - require.NoError(t, err) + require.NoError(t, y.SetBytesCanonical(b)) if !y.Equal(&x) { - t.Fatal("SetBytes(Marshal(x)) != x") + t.Fatal("SetBytesCanonical(Marshal(x)) != x") } } for _, n := range []int{0, BytesE6 - 1, BytesE6 + 1} { var z E6 - if _, err := z.SetBytes(make([]byte, n)); err == nil { - t.Fatalf("SetBytes did not fail on %d bytes", n) + if err := z.SetBytesCanonical(make([]byte, n)); err == nil { + t.Fatalf("SetBytesCanonical did not fail on %d bytes", n) } } } + +func TestE6SetBytesCanonicalRejectsNonCanonical(t *testing.T) { + var x E6 + x.MustSetRandom() + good := x.Marshal() + + q := fr.Modulus() + for i := range 6 { + b := append([]byte(nil), good...) + q.FillBytes(b[i*fr.Bytes : (i+1)*fr.Bytes]) + + y := x + require.Error(t, y.SetBytesCanonical(b), "coefficient %d equal to the modulus must be rejected", i) + require.True(t, y.Equal(&x), "E6 must be unchanged on error") + } +} diff --git a/field/mamabear/extensions/e3.go b/field/mamabear/extensions/e3.go index 61c604abb5..baaac5d57e 100644 --- a/field/mamabear/extensions/e3.go +++ b/field/mamabear/extensions/e3.go @@ -90,16 +90,25 @@ func (z *E3) Marshal() []byte { return res } -// SetBytes sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytes, and returns z. -// It returns an error if len(b) != BytesE3. -func (z *E3) SetBytes(b []byte) (*E3, error) { +// SetBytesCanonical sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytesCanonical. +// It returns an error if len(b) != BytesE3 or if any coefficient is not the canonical encoding of a field element; +// in that case z is left unchanged. +func (z *E3) SetBytesCanonical(b []byte) error { if len(b) != BytesE3 { - return nil, fmt.Errorf("E3.SetBytes: got %d bytes, expected %d", len(b), BytesE3) + return fmt.Errorf("E3.SetBytesCanonical: got %d bytes, expected %d", len(b), BytesE3) } - z.A0.SetBytes(b[0*fr.Bytes : 1*fr.Bytes]) - z.A1.SetBytes(b[1*fr.Bytes : 2*fr.Bytes]) - z.A2.SetBytes(b[2*fr.Bytes : 3*fr.Bytes]) - return z, nil + var r E3 + if err := r.A0.SetBytesCanonical(b[0*fr.Bytes : 1*fr.Bytes]); err != nil { + return err + } + if err := r.A1.SetBytesCanonical(b[1*fr.Bytes : 2*fr.Bytes]); err != nil { + return err + } + if err := r.A2.SetBytesCanonical(b[2*fr.Bytes : 3*fr.Bytes]); err != nil { + return err + } + *z = r + return nil } // IsZero reports whether z is zero. diff --git a/field/mamabear/extensions/e3_test.go b/field/mamabear/extensions/e3_test.go index eda81b7e86..c83c1f1f07 100644 --- a/field/mamabear/extensions/e3_test.go +++ b/field/mamabear/extensions/e3_test.go @@ -446,7 +446,7 @@ func BenchmarkE3VectorMulAccByElement(b *testing.B) { } } -func TestE3MarshalSetBytesRoundTrip(t *testing.T) { +func TestE3MarshalSetBytesCanonicalRoundTrip(t *testing.T) { for range 100 { var x E3 x.MustSetRandom() @@ -464,17 +464,16 @@ func TestE3MarshalSetBytesRoundTrip(t *testing.T) { } var y E3 - _, err := y.SetBytes(b) - require.NoError(t, err) + require.NoError(t, y.SetBytesCanonical(b)) if !y.Equal(&x) { - t.Fatal("SetBytes(Marshal(x)) != x") + t.Fatal("SetBytesCanonical(Marshal(x)) != x") } } for _, n := range []int{0, BytesE3 - 1, BytesE3 + 1} { var z E3 - if _, err := z.SetBytes(make([]byte, n)); err == nil { - t.Fatalf("SetBytes did not fail on %d bytes", n) + if err := z.SetBytesCanonical(make([]byte, n)); err == nil { + t.Fatalf("SetBytesCanonical did not fail on %d bytes", n) } } } @@ -528,3 +527,19 @@ func TestE3Div(t *testing.T) { require.True(t, xc.Equal(&q), "Div must be alias-safe") } } + +func TestE3SetBytesCanonicalRejectsNonCanonical(t *testing.T) { + var x E3 + x.MustSetRandom() + good := x.Marshal() + + q := fr.Modulus() + for i := range 3 { + b := append([]byte(nil), good...) + q.FillBytes(b[i*fr.Bytes : (i+1)*fr.Bytes]) + + y := x + require.Error(t, y.SetBytesCanonical(b), "coefficient %d equal to the modulus must be rejected", i) + require.True(t, y.Equal(&x), "E3 must be unchanged on error") + } +} diff --git a/internal/generator/field/template/extensions/e2.go.tmpl b/internal/generator/field/template/extensions/e2.go.tmpl index a8970fe6cf..4659d35f11 100644 --- a/internal/generator/field/template/extensions/e2.go.tmpl +++ b/internal/generator/field/template/extensions/e2.go.tmpl @@ -106,15 +106,22 @@ func (z *E2) Marshal() []byte { return res } -// SetBytes sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytes, and returns z. -// It returns an error if len(b) != BytesE2. -func (z *E2) SetBytes(b []byte) (*E2, error) { +// SetBytesCanonical sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytesCanonical. +// It returns an error if len(b) != BytesE2 or if any coefficient is not the canonical encoding of a field element; +// in that case z is left unchanged. +func (z *E2) SetBytesCanonical(b []byte) error { if len(b) != BytesE2 { - return nil, fmt.Errorf("E2.SetBytes: got %d bytes, expected %d", len(b), BytesE2) + return fmt.Errorf("E2.SetBytesCanonical: got %d bytes, expected %d", len(b), BytesE2) } - z.A0.SetBytes(b[0*fr.Bytes : 1*fr.Bytes]) - z.A1.SetBytes(b[1*fr.Bytes : 2*fr.Bytes]) - return z, nil + var r E2 + if err := r.A0.SetBytesCanonical(b[0*fr.Bytes : 1*fr.Bytes]); err != nil { + return err + } + if err := r.A1.SetBytesCanonical(b[1*fr.Bytes : 2*fr.Bytes]); err != nil { + return err + } + *z = r + return nil } // SetRandom sets a0 and a1 to random values diff --git a/internal/generator/field/template/extensions/e2_test.go.tmpl b/internal/generator/field/template/extensions/e2_test.go.tmpl index 89c55f9021..84398b3bec 100644 --- a/internal/generator/field/template/extensions/e2_test.go.tmpl +++ b/internal/generator/field/template/extensions/e2_test.go.tmpl @@ -561,7 +561,7 @@ func genE2() gopter.Gen { }) } -func TestE2MarshalSetBytesRoundTrip(t *testing.T) { +func TestE2MarshalSetBytesCanonicalRoundTrip(t *testing.T) { for range 100 { var x E2 x.MustSetRandom() @@ -579,17 +579,16 @@ func TestE2MarshalSetBytesRoundTrip(t *testing.T) { } var y E2 - _, err := y.SetBytes(b) - require.NoError(t, err) + require.NoError(t, y.SetBytesCanonical(b)) if !y.Equal(&x) { - t.Fatal("SetBytes(Marshal(x)) != x") + t.Fatal("SetBytesCanonical(Marshal(x)) != x") } } for _, n := range []int{0, BytesE2 - 1, BytesE2 + 1} { var z E2 - if _, err := z.SetBytes(make([]byte, n)); err == nil { - t.Fatalf("SetBytes did not fail on %d bytes", n) + if err := z.SetBytesCanonical(make([]byte, n)); err == nil { + t.Fatalf("SetBytesCanonical did not fail on %d bytes", n) } } } @@ -621,3 +620,19 @@ func TestE2SetInt64SetUint64(t *testing.T) { require.Truef(t, z.Equal(&zb), "SetUint64(%d) != SetBigInt(%d)", v, v) } } + +func TestE2SetBytesCanonicalRejectsNonCanonical(t *testing.T) { + var x E2 + x.MustSetRandom() + good := x.Marshal() + + q := fr.Modulus() + for i := range 2 { + b := append([]byte(nil), good...) + q.FillBytes(b[i*fr.Bytes : (i+1)*fr.Bytes]) + + y := x + require.Error(t, y.SetBytesCanonical(b), "coefficient %d equal to the modulus must be rejected", i) + require.True(t, y.Equal(&x), "E2 must be unchanged on error") + } +} diff --git a/internal/generator/field/template/extensions/e3.go.tmpl b/internal/generator/field/template/extensions/e3.go.tmpl index f07396b0cb..959ee05184 100644 --- a/internal/generator/field/template/extensions/e3.go.tmpl +++ b/internal/generator/field/template/extensions/e3.go.tmpl @@ -82,16 +82,25 @@ func (z *E3) Marshal() []byte { return res } -// SetBytes sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytes, and returns z. -// It returns an error if len(b) != BytesE3. -func (z *E3) SetBytes(b []byte) (*E3, error) { +// SetBytesCanonical sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytesCanonical. +// It returns an error if len(b) != BytesE3 or if any coefficient is not the canonical encoding of a field element; +// in that case z is left unchanged. +func (z *E3) SetBytesCanonical(b []byte) error { if len(b) != BytesE3 { - return nil, fmt.Errorf("E3.SetBytes: got %d bytes, expected %d", len(b), BytesE3) + return fmt.Errorf("E3.SetBytesCanonical: got %d bytes, expected %d", len(b), BytesE3) } - z.A0.SetBytes(b[0*fr.Bytes : 1*fr.Bytes]) - z.A1.SetBytes(b[1*fr.Bytes : 2*fr.Bytes]) - z.A2.SetBytes(b[2*fr.Bytes : 3*fr.Bytes]) - return z, nil + var r E3 + if err := r.A0.SetBytesCanonical(b[0*fr.Bytes : 1*fr.Bytes]); err != nil { + return err + } + if err := r.A1.SetBytesCanonical(b[1*fr.Bytes : 2*fr.Bytes]); err != nil { + return err + } + if err := r.A2.SetBytesCanonical(b[2*fr.Bytes : 3*fr.Bytes]); err != nil { + return err + } + *z = r + return nil } // IsZero reports whether z is zero. diff --git a/internal/generator/field/template/extensions/e3_test.go.tmpl b/internal/generator/field/template/extensions/e3_test.go.tmpl index 5b0a56113a..83f5fe8cda 100644 --- a/internal/generator/field/template/extensions/e3_test.go.tmpl +++ b/internal/generator/field/template/extensions/e3_test.go.tmpl @@ -437,7 +437,7 @@ func BenchmarkE3VectorMulAccByElement(b *testing.B) { } } -func TestE3MarshalSetBytesRoundTrip(t *testing.T) { +func TestE3MarshalSetBytesCanonicalRoundTrip(t *testing.T) { for range 100 { var x E3 x.MustSetRandom() @@ -455,17 +455,16 @@ func TestE3MarshalSetBytesRoundTrip(t *testing.T) { } var y E3 - _, err := y.SetBytes(b) - require.NoError(t, err) + require.NoError(t, y.SetBytesCanonical(b)) if !y.Equal(&x) { - t.Fatal("SetBytes(Marshal(x)) != x") + t.Fatal("SetBytesCanonical(Marshal(x)) != x") } } for _, n := range []int{0, BytesE3 - 1, BytesE3 + 1} { var z E3 - if _, err := z.SetBytes(make([]byte, n)); err == nil { - t.Fatalf("SetBytes did not fail on %d bytes", n) + if err := z.SetBytesCanonical(make([]byte, n)); err == nil { + t.Fatalf("SetBytesCanonical did not fail on %d bytes", n) } } } @@ -519,3 +518,19 @@ func TestE3Div(t *testing.T) { require.True(t, xc.Equal(&q), "Div must be alias-safe") } } + +func TestE3SetBytesCanonicalRejectsNonCanonical(t *testing.T) { + var x E3 + x.MustSetRandom() + good := x.Marshal() + + q := fr.Modulus() + for i := range 3 { + b := append([]byte(nil), good...) + q.FillBytes(b[i*fr.Bytes : (i+1)*fr.Bytes]) + + y := x + require.Error(t, y.SetBytesCanonical(b), "coefficient %d equal to the modulus must be rejected", i) + require.True(t, y.Equal(&x), "E3 must be unchanged on error") + } +} diff --git a/internal/generator/field/template/extensions/e4.go.tmpl b/internal/generator/field/template/extensions/e4.go.tmpl index 96a77c6bf2..766e58d57c 100644 --- a/internal/generator/field/template/extensions/e4.go.tmpl +++ b/internal/generator/field/template/extensions/e4.go.tmpl @@ -119,17 +119,28 @@ func (z *E4) Marshal() []byte { return res } -// SetBytes sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytes, and returns z. -// It returns an error if len(b) != BytesE4. -func (z *E4) SetBytes(b []byte) (*E4, error) { +// SetBytesCanonical sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytesCanonical. +// It returns an error if len(b) != BytesE4 or if any coefficient is not the canonical encoding of a field element; +// in that case z is left unchanged. +func (z *E4) SetBytesCanonical(b []byte) error { if len(b) != BytesE4 { - return nil, fmt.Errorf("E4.SetBytes: got %d bytes, expected %d", len(b), BytesE4) + return fmt.Errorf("E4.SetBytesCanonical: got %d bytes, expected %d", len(b), BytesE4) } - z.B0.A0.SetBytes(b[0*fr.Bytes : 1*fr.Bytes]) - z.B0.A1.SetBytes(b[1*fr.Bytes : 2*fr.Bytes]) - z.B1.A0.SetBytes(b[2*fr.Bytes : 3*fr.Bytes]) - z.B1.A1.SetBytes(b[3*fr.Bytes : 4*fr.Bytes]) - return z, nil + var r E4 + if err := r.B0.A0.SetBytesCanonical(b[0*fr.Bytes : 1*fr.Bytes]); err != nil { + return err + } + if err := r.B0.A1.SetBytesCanonical(b[1*fr.Bytes : 2*fr.Bytes]); err != nil { + return err + } + if err := r.B1.A0.SetBytesCanonical(b[2*fr.Bytes : 3*fr.Bytes]); err != nil { + return err + } + if err := r.B1.A1.SetBytesCanonical(b[3*fr.Bytes : 4*fr.Bytes]); err != nil { + return err + } + *z = r + return nil } // Lift sets the B0.A0 component of z to v diff --git a/internal/generator/field/template/extensions/e4_test.go.tmpl b/internal/generator/field/template/extensions/e4_test.go.tmpl index 968f4662b4..dcf32e9daf 100644 --- a/internal/generator/field/template/extensions/e4_test.go.tmpl +++ b/internal/generator/field/template/extensions/e4_test.go.tmpl @@ -1125,7 +1125,7 @@ func genFrVector(size int) gopter.Gen { } } -func TestE4MarshalSetBytesRoundTrip(t *testing.T) { +func TestE4MarshalSetBytesCanonicalRoundTrip(t *testing.T) { for range 100 { var x E4 x.MustSetRandom() @@ -1143,17 +1143,16 @@ func TestE4MarshalSetBytesRoundTrip(t *testing.T) { } var y E4 - _, err := y.SetBytes(b) - require.NoError(t, err) + require.NoError(t, y.SetBytesCanonical(b)) if !y.Equal(&x) { - t.Fatal("SetBytes(Marshal(x)) != x") + t.Fatal("SetBytesCanonical(Marshal(x)) != x") } } for _, n := range []int{0, BytesE4 - 1, BytesE4 + 1} { var z E4 - if _, err := z.SetBytes(make([]byte, n)); err == nil { - t.Fatalf("SetBytes did not fail on %d bytes", n) + if err := z.SetBytesCanonical(make([]byte, n)); err == nil { + t.Fatalf("SetBytesCanonical did not fail on %d bytes", n) } } } @@ -1189,3 +1188,19 @@ func TestE4SetInt64SetUint64(t *testing.T) { require.Truef(t, z.Equal(&zb), "SetUint64(%d) != SetBigInt(%d)", v, v) } } + +func TestE4SetBytesCanonicalRejectsNonCanonical(t *testing.T) { + var x E4 + x.MustSetRandom() + good := x.Marshal() + + q := fr.Modulus() + for i := range 4 { + b := append([]byte(nil), good...) + q.FillBytes(b[i*fr.Bytes : (i+1)*fr.Bytes]) + + y := x + require.Error(t, y.SetBytesCanonical(b), "coefficient %d equal to the modulus must be rejected", i) + require.True(t, y.Equal(&x), "E4 must be unchanged on error") + } +} diff --git a/internal/generator/field/template/extensions/e6.go.tmpl b/internal/generator/field/template/extensions/e6.go.tmpl index c028861084..7f996bbcc2 100644 --- a/internal/generator/field/template/extensions/e6.go.tmpl +++ b/internal/generator/field/template/extensions/e6.go.tmpl @@ -114,19 +114,34 @@ func (z *E6) Marshal() []byte { return res } -// SetBytes sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytes, and returns z. -// It returns an error if len(b) != BytesE6. -func (z *E6) SetBytes(b []byte) (*E6, error) { +// SetBytesCanonical sets z from the layout produced by Marshal, reading each coefficient with fr.Element.SetBytesCanonical. +// It returns an error if len(b) != BytesE6 or if any coefficient is not the canonical encoding of a field element; +// in that case z is left unchanged. +func (z *E6) SetBytesCanonical(b []byte) error { if len(b) != BytesE6 { - return nil, fmt.Errorf("E6.SetBytes: got %d bytes, expected %d", len(b), BytesE6) - } - z.B0.A0.SetBytes(b[0*fr.Bytes : 1*fr.Bytes]) - z.B0.A1.SetBytes(b[1*fr.Bytes : 2*fr.Bytes]) - z.B1.A0.SetBytes(b[2*fr.Bytes : 3*fr.Bytes]) - z.B1.A1.SetBytes(b[3*fr.Bytes : 4*fr.Bytes]) - z.B2.A0.SetBytes(b[4*fr.Bytes : 5*fr.Bytes]) - z.B2.A1.SetBytes(b[5*fr.Bytes : 6*fr.Bytes]) - return z, nil + return fmt.Errorf("E6.SetBytesCanonical: got %d bytes, expected %d", len(b), BytesE6) + } + var r E6 + if err := r.B0.A0.SetBytesCanonical(b[0*fr.Bytes : 1*fr.Bytes]); err != nil { + return err + } + if err := r.B0.A1.SetBytesCanonical(b[1*fr.Bytes : 2*fr.Bytes]); err != nil { + return err + } + if err := r.B1.A0.SetBytesCanonical(b[2*fr.Bytes : 3*fr.Bytes]); err != nil { + return err + } + if err := r.B1.A1.SetBytesCanonical(b[3*fr.Bytes : 4*fr.Bytes]); err != nil { + return err + } + if err := r.B2.A0.SetBytesCanonical(b[4*fr.Bytes : 5*fr.Bytes]); err != nil { + return err + } + if err := r.B2.A1.SetBytesCanonical(b[5*fr.Bytes : 6*fr.Bytes]); err != nil { + return err + } + *z = r + return nil } // MulByElement multiplies an element in E6 by an element in fr. diff --git a/internal/generator/field/template/extensions/e6_test.go.tmpl b/internal/generator/field/template/extensions/e6_test.go.tmpl index 752fe6aabb..d6c4fdcadc 100644 --- a/internal/generator/field/template/extensions/e6_test.go.tmpl +++ b/internal/generator/field/template/extensions/e6_test.go.tmpl @@ -352,7 +352,7 @@ func genE6() gopter.Gen { }) } -func TestE6MarshalSetBytesRoundTrip(t *testing.T) { +func TestE6MarshalSetBytesCanonicalRoundTrip(t *testing.T) { for range 100 { var x E6 x.MustSetRandom() @@ -370,17 +370,32 @@ func TestE6MarshalSetBytesRoundTrip(t *testing.T) { } var y E6 - _, err := y.SetBytes(b) - require.NoError(t, err) + require.NoError(t, y.SetBytesCanonical(b)) if !y.Equal(&x) { - t.Fatal("SetBytes(Marshal(x)) != x") + t.Fatal("SetBytesCanonical(Marshal(x)) != x") } } for _, n := range []int{0, BytesE6 - 1, BytesE6 + 1} { var z E6 - if _, err := z.SetBytes(make([]byte, n)); err == nil { - t.Fatalf("SetBytes did not fail on %d bytes", n) + if err := z.SetBytesCanonical(make([]byte, n)); err == nil { + t.Fatalf("SetBytesCanonical did not fail on %d bytes", n) } } } + +func TestE6SetBytesCanonicalRejectsNonCanonical(t *testing.T) { + var x E6 + x.MustSetRandom() + good := x.Marshal() + + q := fr.Modulus() + for i := range 6 { + b := append([]byte(nil), good...) + q.FillBytes(b[i*fr.Bytes : (i+1)*fr.Bytes]) + + y := x + require.Error(t, y.SetBytesCanonical(b), "coefficient %d equal to the modulus must be rejected", i) + require.True(t, y.Equal(&x), "E6 must be unchanged on error") + } +} From 27a80de95fadbd01e0c859f9b66af280803f0925 Mon Sep 17 00:00:00 2001 From: Arya Tabaie Date: Mon, 5 Oct 2026 15:44:01 -0500 Subject: [PATCH 10/18] feat: mixed operations Signed-off-by: Arya Tabaie --- field/babybear/extensions/e2.go | 24 ++ field/babybear/extensions/e2_test.go | 31 ++ field/babybear/extensions/e4.go | 28 ++ field/babybear/extensions/e4_test.go | 33 ++ field/babybear/extensions/e6.go | 32 ++ field/babybear/extensions/e6_test.go | 35 ++ field/babybear/polynomial/doc.go | 7 + field/babybear/polynomial/multilin.go | 179 ++++++++++ field/babybear/polynomial/multilin_test.go | 85 +++++ field/babybear/polynomial/polynomial.go | 310 ++++++++++++++++++ field/babybear/polynomial/polynomial_test.go | 255 ++++++++++++++ field/babybear/polynomial/pool.go | 191 +++++++++++ field/goldilocks/extensions/e2.go | 24 ++ field/goldilocks/extensions/e2_test.go | 31 ++ field/goldilocks/polynomial/doc.go | 7 + field/goldilocks/polynomial/multilin.go | 179 ++++++++++ field/goldilocks/polynomial/multilin_test.go | 85 +++++ field/goldilocks/polynomial/polynomial.go | 310 ++++++++++++++++++ .../goldilocks/polynomial/polynomial_test.go | 255 ++++++++++++++ field/goldilocks/polynomial/pool.go | 191 +++++++++++ field/koalabear/extensions/e2.go | 24 ++ field/koalabear/extensions/e2_test.go | 31 ++ field/koalabear/extensions/e4.go | 28 ++ field/koalabear/extensions/e4_test.go | 33 ++ field/koalabear/extensions/e6.go | 32 ++ field/koalabear/extensions/e6_test.go | 35 ++ .../extensions/polynomial/multilin_e6.go | 63 ++++ .../extensions/polynomial/multilin_e6_test.go | 63 ++++ field/koalabear/polynomial/doc.go | 7 + field/koalabear/polynomial/multilin.go | 179 ++++++++++ field/koalabear/polynomial/multilin_test.go | 85 +++++ field/koalabear/polynomial/polynomial.go | 310 ++++++++++++++++++ field/koalabear/polynomial/polynomial_test.go | 255 ++++++++++++++ field/koalabear/polynomial/pool.go | 191 +++++++++++ field/mamabear/extensions/e3.go | 26 ++ field/mamabear/extensions/e3_test.go | 32 ++ field/mamabear/polynomial/doc.go | 7 + field/mamabear/polynomial/multilin.go | 179 ++++++++++ field/mamabear/polynomial/multilin_test.go | 85 +++++ field/mamabear/polynomial/polynomial.go | 310 ++++++++++++++++++ field/mamabear/polynomial/polynomial_test.go | 255 ++++++++++++++ field/mamabear/polynomial/pool.go | 191 +++++++++++ .../generator/field/config/field_config.go | 2 + .../field/template/extensions/e2.go.tmpl | 24 ++ .../field/template/extensions/e2_test.go.tmpl | 31 ++ .../field/template/extensions/e3.go.tmpl | 26 ++ .../field/template/extensions/e3_test.go.tmpl | 32 ++ .../field/template/extensions/e4.go.tmpl | 28 ++ .../field/template/extensions/e4_test.go.tmpl | 33 ++ .../field/template/extensions/e6.go.tmpl | 32 ++ .../field/template/extensions/e6_test.go.tmpl | 35 ++ internal/generator/main.go | 16 +- .../polynomial/template/multilin.go.tmpl | 67 ++++ .../polynomial/template/multilin.test.go.tmpl | 69 ++++ 54 files changed, 5104 insertions(+), 4 deletions(-) create mode 100644 field/babybear/polynomial/doc.go create mode 100644 field/babybear/polynomial/multilin.go create mode 100644 field/babybear/polynomial/multilin_test.go create mode 100644 field/babybear/polynomial/polynomial.go create mode 100644 field/babybear/polynomial/polynomial_test.go create mode 100644 field/babybear/polynomial/pool.go create mode 100644 field/goldilocks/polynomial/doc.go create mode 100644 field/goldilocks/polynomial/multilin.go create mode 100644 field/goldilocks/polynomial/multilin_test.go create mode 100644 field/goldilocks/polynomial/polynomial.go create mode 100644 field/goldilocks/polynomial/polynomial_test.go create mode 100644 field/goldilocks/polynomial/pool.go create mode 100644 field/koalabear/polynomial/doc.go create mode 100644 field/koalabear/polynomial/multilin.go create mode 100644 field/koalabear/polynomial/multilin_test.go create mode 100644 field/koalabear/polynomial/polynomial.go create mode 100644 field/koalabear/polynomial/polynomial_test.go create mode 100644 field/koalabear/polynomial/pool.go create mode 100644 field/mamabear/polynomial/doc.go create mode 100644 field/mamabear/polynomial/multilin.go create mode 100644 field/mamabear/polynomial/multilin_test.go create mode 100644 field/mamabear/polynomial/polynomial.go create mode 100644 field/mamabear/polynomial/polynomial_test.go create mode 100644 field/mamabear/polynomial/pool.go diff --git a/field/babybear/extensions/e2.go b/field/babybear/extensions/e2.go index 56371d4912..ae19692e74 100644 --- a/field/babybear/extensions/e2.go +++ b/field/babybear/extensions/e2.go @@ -89,6 +89,30 @@ func (z *E2) SetUint64(v uint64) *E2 { return z } +// AddElement sets z to x + y, where y is an element of the base field embedded in E2, and returns z +func (z *E2) AddElement(x *E2, y *fr.Element) *E2 { + yc := *y + z.A0.Add(&x.A0, &yc) + z.A1 = x.A1 + return z +} + +// SubElement sets z to x - y, where y is an element of the base field embedded in E2, and returns z +func (z *E2) SubElement(x *E2, y *fr.Element) *E2 { + yc := *y + z.A0.Sub(&x.A0, &yc) + z.A1 = x.A1 + return z +} + +// SetElement sets z to x, an element of the base field embedded in E2, and returns z +func (z *E2) SetElement(x *fr.Element) *E2 { + v := *x + *z = E2{} + z.A0 = v + return z +} + // SetBigInt sets z to v (embedded via the unique ring homomorphism ℤ→𝔽, reduced modulo the base field order) and returns z func (z *E2) SetBigInt(v *big.Int) *E2 { *z = E2{} diff --git a/field/babybear/extensions/e2_test.go b/field/babybear/extensions/e2_test.go index 37e3d7d4c1..a649b3c495 100644 --- a/field/babybear/extensions/e2_test.go +++ b/field/babybear/extensions/e2_test.go @@ -632,3 +632,34 @@ func TestE2SetBytesCanonicalRejectsNonCanonical(t *testing.T) { require.True(t, y.Equal(&x), "E2 must be unchanged on error") } } + +func TestE2ElementOps(t *testing.T) { + for range 100 { + var x, lifted, got, want E2 + var e fr.Element + x.MustSetRandom() + e.MustSetRandom() + + var z E2 + z.SetElement(&e) + require.True(t, z.A0.Equal(&e)) + require.True(t, z.A1.IsZero()) + lifted = z + + got.AddElement(&x, &e) + want.Add(&x, &lifted) + require.True(t, got.Equal(&want), "AddElement") + + got.SubElement(&x, &e) + want.Sub(&x, &lifted) + require.True(t, got.Equal(&want), "SubElement") + + // aliasing of the receiver with the first operand + got.Set(&x) + got.AddElement(&got, &e) + want.Add(&x, &lifted) + require.True(t, got.Equal(&want), "AddElement alias") + got.SubElement(&got, &e) + require.True(t, got.Equal(&x), "SubElement alias") + } +} diff --git a/field/babybear/extensions/e4.go b/field/babybear/extensions/e4.go index 879dc452f6..3cab206a80 100644 --- a/field/babybear/extensions/e4.go +++ b/field/babybear/extensions/e4.go @@ -100,6 +100,34 @@ func (z *E4) SetUint64(v uint64) *E4 { return z } +// AddElement sets z to x + y, where y is an element of the base field embedded in E4, and returns z +func (z *E4) AddElement(x *E4, y *fr.Element) *E4 { + yc := *y + z.B0.A0.Add(&x.B0.A0, &yc) + z.B0.A1 = x.B0.A1 + z.B1.A0 = x.B1.A0 + z.B1.A1 = x.B1.A1 + return z +} + +// SubElement sets z to x - y, where y is an element of the base field embedded in E4, and returns z +func (z *E4) SubElement(x *E4, y *fr.Element) *E4 { + yc := *y + z.B0.A0.Sub(&x.B0.A0, &yc) + z.B0.A1 = x.B0.A1 + z.B1.A0 = x.B1.A0 + z.B1.A1 = x.B1.A1 + return z +} + +// SetElement sets z to x, an element of the base field embedded in E4, and returns z +func (z *E4) SetElement(x *fr.Element) *E4 { + v := *x + *z = E4{} + z.B0.A0 = v + return z +} + // SetBigInt sets z to v (embedded via the unique ring homomorphism ℤ→𝔽, reduced modulo the base field order) and returns z func (z *E4) SetBigInt(v *big.Int) *E4 { *z = E4{} diff --git a/field/babybear/extensions/e4_test.go b/field/babybear/extensions/e4_test.go index 2e514ab2dc..2c6189b849 100644 --- a/field/babybear/extensions/e4_test.go +++ b/field/babybear/extensions/e4_test.go @@ -1206,3 +1206,36 @@ func TestE4SetBytesCanonicalRejectsNonCanonical(t *testing.T) { require.True(t, y.Equal(&x), "E4 must be unchanged on error") } } + +func TestE4ElementOps(t *testing.T) { + for range 100 { + var x, lifted, got, want E4 + var e fr.Element + x.MustSetRandom() + e.MustSetRandom() + + var z E4 + z.SetElement(&e) + require.True(t, z.B0.A0.Equal(&e)) + require.True(t, z.B0.A1.IsZero()) + require.True(t, z.B1.A0.IsZero()) + require.True(t, z.B1.A1.IsZero()) + lifted = z + + got.AddElement(&x, &e) + want.Add(&x, &lifted) + require.True(t, got.Equal(&want), "AddElement") + + got.SubElement(&x, &e) + want.Sub(&x, &lifted) + require.True(t, got.Equal(&want), "SubElement") + + // aliasing of the receiver with the first operand + got.Set(&x) + got.AddElement(&got, &e) + want.Add(&x, &lifted) + require.True(t, got.Equal(&want), "AddElement alias") + got.SubElement(&got, &e) + require.True(t, got.Equal(&x), "SubElement alias") + } +} diff --git a/field/babybear/extensions/e6.go b/field/babybear/extensions/e6.go index f9569184a4..f37f878415 100644 --- a/field/babybear/extensions/e6.go +++ b/field/babybear/extensions/e6.go @@ -90,6 +90,38 @@ func (z *E6) SetUint64(v uint64) *E6 { return z } +// AddElement sets z to x + y, where y is an element of the base field embedded in E6, and returns z +func (z *E6) AddElement(x *E6, y *fr.Element) *E6 { + yc := *y + z.B0.A0.Add(&x.B0.A0, &yc) + z.B0.A1 = x.B0.A1 + z.B1.A0 = x.B1.A0 + z.B1.A1 = x.B1.A1 + z.B2.A0 = x.B2.A0 + z.B2.A1 = x.B2.A1 + return z +} + +// SubElement sets z to x - y, where y is an element of the base field embedded in E6, and returns z +func (z *E6) SubElement(x *E6, y *fr.Element) *E6 { + yc := *y + z.B0.A0.Sub(&x.B0.A0, &yc) + z.B0.A1 = x.B0.A1 + z.B1.A0 = x.B1.A0 + z.B1.A1 = x.B1.A1 + z.B2.A0 = x.B2.A0 + z.B2.A1 = x.B2.A1 + return z +} + +// SetElement sets z to x, an element of the base field embedded in E6, and returns z +func (z *E6) SetElement(x *fr.Element) *E6 { + v := *x + *z = E6{} + z.B0.A0 = v + return z +} + // SetBigInt sets z to v (embedded via the unique ring homomorphism ℤ→𝔽, reduced modulo the base field order) and returns z func (z *E6) SetBigInt(v *big.Int) *E6 { *z = E6{} diff --git a/field/babybear/extensions/e6_test.go b/field/babybear/extensions/e6_test.go index 91a0a10d0b..d13acbbf66 100644 --- a/field/babybear/extensions/e6_test.go +++ b/field/babybear/extensions/e6_test.go @@ -409,3 +409,38 @@ func TestE6SetBytesCanonicalRejectsNonCanonical(t *testing.T) { require.True(t, y.Equal(&x), "E6 must be unchanged on error") } } + +func TestE6ElementOps(t *testing.T) { + for range 100 { + var x, lifted, got, want E6 + var e fr.Element + x.MustSetRandom() + e.MustSetRandom() + + var z E6 + z.SetElement(&e) + require.True(t, z.B0.A0.Equal(&e)) + require.True(t, z.B0.A1.IsZero()) + require.True(t, z.B1.A0.IsZero()) + require.True(t, z.B1.A1.IsZero()) + require.True(t, z.B2.A0.IsZero()) + require.True(t, z.B2.A1.IsZero()) + lifted = z + + got.AddElement(&x, &e) + want.Add(&x, &lifted) + require.True(t, got.Equal(&want), "AddElement") + + got.SubElement(&x, &e) + want.Sub(&x, &lifted) + require.True(t, got.Equal(&want), "SubElement") + + // aliasing of the receiver with the first operand + got.Set(&x) + got.AddElement(&got, &e) + want.Add(&x, &lifted) + require.True(t, got.Equal(&want), "AddElement alias") + got.SubElement(&got, &e) + require.True(t, got.Equal(&x), "SubElement alias") + } +} diff --git a/field/babybear/polynomial/doc.go b/field/babybear/polynomial/doc.go new file mode 100644 index 0000000000..aa346f3ea3 --- /dev/null +++ b/field/babybear/polynomial/doc.go @@ -0,0 +1,7 @@ +// Copyright 2020-2026 Consensys Software Inc. +// Licensed under the Apache License, Version 2.0. See the LICENSE file for details. + +// Code generated by consensys/gnark-crypto DO NOT EDIT + +// Package polynomial provides polynomial methods and commitment schemes. +package polynomial diff --git a/field/babybear/polynomial/multilin.go b/field/babybear/polynomial/multilin.go new file mode 100644 index 0000000000..f29a16a4fc --- /dev/null +++ b/field/babybear/polynomial/multilin.go @@ -0,0 +1,179 @@ +// Copyright 2020-2026 Consensys Software Inc. +// Licensed under the Apache License, Version 2.0. See the LICENSE file for details. + +// Code generated by consensys/gnark-crypto DO NOT EDIT + +package polynomial + +import ( + "math/bits" + + "github.com/consensys/gnark-crypto/field/babybear" + "github.com/consensys/gnark-crypto/utils" +) + +// MultiLin tracks the values of a (dense i.e. not sparse) multilinear polynomial +// The variables are X₁ through Xₙ where n = log(len(.)) +// .[∑ᵢ 2ⁱ⁻¹ bₙ₋ᵢ] = the polynomial evaluated at (b₁, b₂, ..., bₙ) +// It is understood that any hypercube evaluation can be extrapolated to a multilinear polynomial +type MultiLin []babybear.Element + +// Fold is partial evaluation function k[X₁, X₂, ..., Xₙ] → k[X₂, ..., Xₙ] by setting X₁=r +func (m *MultiLin) Fold(r babybear.Element) { + mid := len(*m) / 2 + + bottom, top := (*m)[:mid], (*m)[mid:] + + var t babybear.Element // no need to update the top part + + // updating bookkeeping table + // knowing that the polynomial f ∈ (k[X₂, ..., Xₙ])[X₁] is linear, we would get f(r) = f(0) + r(f(1) - f(0)) + // the following loop computes the evaluations of f(r) accordingly: + // f(r, b₂, ..., bₙ) = f(0, b₂, ..., bₙ) + r(f(1, b₂, ..., bₙ) - f(0, b₂, ..., bₙ)) + for i := range mid { + // table[i] ← table[i] + r (table[i + mid] - table[i]) + t.Sub(&top[i], &bottom[i]) + t.Mul(&t, &r) + bottom[i].Add(&bottom[i], &t) + } + + *m = (*m)[:mid] +} + +func (m *MultiLin) FoldParallel(r babybear.Element) utils.Task { + mid := len(*m) / 2 + bottom, top := (*m)[:mid], (*m)[mid:] + + *m = bottom + + return func(start, end int) { + var t babybear.Element // no need to update the top part + for i := start; i < end; i++ { + // table[i] ← table[i] + r (table[i + mid] - table[i]) + t.Sub(&top[i], &bottom[i]) + t.Mul(&t, &r) + bottom[i].Add(&bottom[i], &t) + } + } +} + +func (m MultiLin) Sum() babybear.Element { + s := m[0] + for i := 1; i < len(m); i++ { + s.Add(&s, &m[i]) + } + return s +} + +func _clone(m MultiLin, p *Pool) MultiLin { + if p == nil { + return m.Clone() + } else { + return p.Clone(m) + } +} + +func _dump(m MultiLin, p *Pool) { + if p != nil { + p.Dump(m) + } +} + +// Evaluate extrapolate the value of the multilinear polynomial corresponding to m +// on the given coordinates +func (m MultiLin) Evaluate(coordinates []babybear.Element, p *Pool) babybear.Element { + // Folding is a mutating operation + bkCopy := _clone(m, p) + + // Evaluate step by step through repeated folding (i.e. evaluation at the first remaining variable) + for _, r := range coordinates { + bkCopy.Fold(r) + } + + result := bkCopy[0] + + _dump(bkCopy, p) + return result +} + +// Clone creates a deep copy of a bookkeeping table. +// Both multilinear interpolation and sumcheck require folding an underlying +// array, but folding changes the array. To do both one requires a deep copy +// of the bookkeeping table. +func (m MultiLin) Clone() MultiLin { + res := make(MultiLin, len(m)) + copy(res, m) + return res +} + +// Add two bookKeepingTables +func (m *MultiLin) Add(left, right MultiLin) { + size := len(left) + // Check that left and right have the same size + if len(right) != size || len(*m) != size { + panic("left, right and destination must have the right size") + } + + // Add elementwise + for i := range size { + (*m)[i].Add(&left[i], &right[i]) + } +} + +// EvalEq computes Eq(q₁, ... , qₙ, h₁, ... , hₙ) = Π₁ⁿ Eq(qᵢ, hᵢ) +// where Eq(x,y) = xy + (1-x)(1-y) = 1 - x - y + xy + xy interpolates +// +// _________________ +// | | | +// | 0 | 1 | +// |_______|_______| +// y | | | +// | 1 | 0 | +// |_______|_______| +// +// x +// +// In other words the polynomial evaluated here is the multilinear extrapolation of +// one that evaluates to q' == h' for vectors q', h' of binary values +func EvalEq(q, h []babybear.Element) babybear.Element { + var res, nxt, one, sum babybear.Element + one.SetOne() + for i := range len(q) { + nxt.Mul(&q[i], &h[i]) // nxt <- qᵢ * hᵢ + nxt.Double(&nxt) // nxt <- 2 * qᵢ * hᵢ + nxt.Add(&nxt, &one) // nxt <- 1 + 2 * qᵢ * hᵢ + sum.Add(&q[i], &h[i]) // sum <- qᵢ + hᵢ TODO: Why not subtract one by one from nxt? More parallel? + + if i == 0 { + res.Sub(&nxt, &sum) // nxt <- 1 + 2 * qᵢ * hᵢ - qᵢ - hᵢ + } else { + nxt.Sub(&nxt, &sum) // nxt <- 1 + 2 * qᵢ * hᵢ - qᵢ - hᵢ + res.Mul(&res, &nxt) // res <- res * nxt + } + } + return res +} + +// Eq sets m to the representation of the polynomial Eq(q₁, ..., qₙ, *, ..., *) × m[0] +func (m *MultiLin) Eq(q []babybear.Element) { + n := len(q) + + if len(*m) != 1<= 0; i-- { + res.Mul(&res, v) + res.Add(&res, &(*p)[i]) + } + + return res +} + +// Clone returns a copy of the polynomial +func (p *Polynomial) Clone() Polynomial { + _p := make(Polynomial, len(*p)) + copy(_p, *p) + return _p +} + +// Set to another polynomial +func (p *Polynomial) Set(p1 Polynomial) { + if len(*p) != len(p1) { + *p = p1.Clone() + return + } + + for i := range len(p1) { + (*p)[i].Set(&p1[i]) + } +} + +// AddConstantInPlace adds a constant to the polynomial, modifying p +func (p *Polynomial) AddConstantInPlace(c *babybear.Element) { + for i := range len(*p) { + (*p)[i].Add(&(*p)[i], c) + } +} + +// SubConstantInPlace subs a constant to the polynomial, modifying p +func (p *Polynomial) SubConstantInPlace(c *babybear.Element) { + for i := range len(*p) { + (*p)[i].Sub(&(*p)[i], c) + } +} + +// ScaleInPlace multiplies p by v, modifying p +func (p *Polynomial) ScaleInPlace(c *babybear.Element) { + for i := range len(*p) { + (*p)[i].Mul(&(*p)[i], c) + } +} + +// Scale multiplies p0 by v, storing the result in p +func (p *Polynomial) Scale(c *babybear.Element, p0 Polynomial) { + if len(*p) != len(p0) { + *p = make(Polynomial, len(p0)) + } + for i := range len(p0) { + (*p)[i].Mul(c, &p0[i]) + } +} + +// Add adds p1 to p2 +// This function allocates a new slice unless p == p1 or p == p2 +func (p *Polynomial) Add(p1, p2 Polynomial) *Polynomial { + + bigger := p1 + smaller := p2 + if len(bigger) < len(smaller) { + bigger, smaller = smaller, bigger + } + + if len(*p) == len(bigger) && (&(*p)[0] == &bigger[0]) { + for i := range len(smaller) { + (*p)[i].Add(&(*p)[i], &smaller[i]) + } + return p + } + + if len(*p) == len(smaller) && (&(*p)[0] == &smaller[0]) { + for i := range len(smaller) { + (*p)[i].Add(&(*p)[i], &bigger[i]) + } + *p = append(*p, bigger[len(smaller):]...) + return p + } + + res := make(Polynomial, len(bigger)) + copy(res, bigger) + for i := range len(smaller) { + res[i].Add(&res[i], &smaller[i]) + } + *p = res + return p +} + +// Sub subtracts p2 from p1 +// TODO make interface more consistent with Add +func (p *Polynomial) Sub(p1, p2 Polynomial) *Polynomial { + if len(p1) != len(p2) || len(p2) != len(*p) { + return nil + } + for i := range len(*p) { + (*p)[i].Sub(&p1[i], &p2[i]) + } + return p +} + +// Equal checks equality between two polynomials +func (p *Polynomial) Equal(p1 Polynomial) bool { + if (*p == nil) != (p1 == nil) { + return false + } + + if len(*p) != len(p1) { + return false + } + + for i := range p1 { + if !(*p)[i].Equal(&p1[i]) { + return false + } + } + + return true +} + +func (p Polynomial) SetZero() { + for i := range len(p) { + p[i].SetZero() + } +} + +func (p Polynomial) Text(base int) string { + + var builder strings.Builder + + first := true + for d := len(p) - 1; d >= 0; d-- { + if p[d].IsZero() { + continue + } + + pD := p[d] + pDText := pD.Text(base) + + initialLen := builder.Len() + + if pDText[0] == '-' { + pDText = pDText[1:] + if first { + builder.WriteString("-") + } else { + builder.WriteString(" - ") + } + } else if !first { + builder.WriteString(" + ") + } + + first = false + + if !pD.IsOne() || d == 0 { + builder.WriteString(pDText) + } + + if builder.Len()-initialLen > 10 { + builder.WriteString("×") + } + + if d != 0 { + builder.WriteString("X") + } + if d > 1 { + builder.WriteString( + utils.ToSuperscript(strconv.Itoa(d)), + ) + } + + } + + if first { + return "0" + } + + return builder.String() +} + +// InterpolateOnRange maps vector v to polynomial f +// such that f(i) = v[i] for 0 ≤ i < len(v). +// len(f) = len(v) and deg(f) ≤ len(v) - 1 +func InterpolateOnRange(v []babybear.Element) Polynomial { + nEvals := uint8(len(v)) + if int(nEvals) != len(v) { + panic("interpolation method too inefficient for nEvals > 255") + } + lagrange := getLagrangeBasis(nEvals) + + var res Polynomial + res.Scale(&v[0], lagrange[0]) + + temp := make(Polynomial, nEvals) + + for i := uint8(1); i < nEvals; i++ { + temp.Scale(&v[i], lagrange[i]) + res.Add(res, temp) + } + + return res +} + +// lagrange bases used by InterpolateOnRange +var lagrangeBasis sync.Map + +func getLagrangeBasis(domainSize uint8) []Polynomial { + if res, ok := lagrangeBasis.Load(domainSize); ok { + return res.([]Polynomial) + } + + // not found. compute + var res []Polynomial + if domainSize >= 2 { + res = computeLagrangeBasis(domainSize) + } else if domainSize == 1 { + res = []Polynomial{make(Polynomial, 1)} + res[0][0].SetOne() + } + lagrangeBasis.Store(domainSize, res) + + return res +} + +// computeLagrangeBasis precomputes in explicit coefficient form for each 0 ≤ l < domainSize the polynomial +// pₗ := X (X-1) ... (X-l-1) (X-l+1) ... (X - domainSize + 1) / ( l (l-1) ... 2 (-1) ... (l - domainSize +1) ) +// Note that pₗ(l) = 1 and pₗ(n) = 0 if 0 ≤ l < domainSize, n ≠ l +func computeLagrangeBasis(domainSize uint8) []Polynomial { + + constTerms := make([]babybear.Element, domainSize) + for i := range domainSize { + constTerms[i].SetInt64(-int64(i)) + } + + res := make([]Polynomial, domainSize) + multScratch := make(Polynomial, domainSize-1) + + // compute pₗ + for l := range domainSize { + + // TODO @Tabaie Optimize this with some trees? O(log(domainSize)) polynomial mults instead of O(domainSize)? Then again it would be fewer big poly mults vs many small poly mults + d := uint8(0) //d is the current degree of res + for i := range domainSize { + if i == l { + continue + } + if d == 0 { + res[l] = make(Polynomial, domainSize) + res[l][domainSize-2] = constTerms[i] + res[l][domainSize-1].SetOne() + } else { + current := res[l][domainSize-d-2:] + timesConst := multScratch[domainSize-d-2:] + + timesConst.Scale(&constTerms[i], current[1:]) //TODO: Directly double and add since constTerms are tiny? (even less than 4 bits) + nonLeading := current[0 : d+1] + + nonLeading.Add(nonLeading, timesConst) + + } + d++ + } + + } + + // We have pₗ(i≠l)=0. Now scale so that pₗ(l)=1 + // Replace the constTerms with norms + for l := range domainSize { + constTerms[l].Neg(&constTerms[l]) + constTerms[l] = res[l].Eval(&constTerms[l]) + } + constTerms = babybear.BatchInvert(constTerms) + for l := range domainSize { + res[l].ScaleInPlace(&constTerms[l]) + } + + return res +} diff --git a/field/babybear/polynomial/polynomial_test.go b/field/babybear/polynomial/polynomial_test.go new file mode 100644 index 0000000000..efcac7a903 --- /dev/null +++ b/field/babybear/polynomial/polynomial_test.go @@ -0,0 +1,255 @@ +// Copyright 2020-2026 Consensys Software Inc. +// Licensed under the Apache License, Version 2.0. See the LICENSE file for details. + +// Code generated by consensys/gnark-crypto DO NOT EDIT + +package polynomial + +import ( + "math/big" + "testing" + + "github.com/consensys/gnark-crypto/field/babybear" + "github.com/leanovate/gopter" + "github.com/leanovate/gopter/gen" + "github.com/leanovate/gopter/prop" + "github.com/stretchr/testify/assert" +) + +func TestPolynomialEval(t *testing.T) { + + // build polynomial + f := make(Polynomial, 20) + for i := range 20 { + f[i].SetOne() + } + + // random value + var point babybear.Element + point.MustSetRandom() + + // compute manually f(val) + var expectedEval, one, den babybear.Element + var expo big.Int + one.SetOne() + expo.SetUint64(20) + expectedEval.Exp(point, &expo). + Sub(&expectedEval, &one) + den.Sub(&point, &one) + expectedEval.Div(&expectedEval, &den) + + // compute purported evaluation + purportedEval := f.Eval(&point) + + // check + if !purportedEval.Equal(&expectedEval) { + t.Fatal("polynomial evaluation failed") + } +} + +func TestPolynomialAddConstantInPlace(t *testing.T) { + + // build polynomial + f := make(Polynomial, 20) + for i := range 20 { + f[i].SetOne() + } + + // constant to add + var c babybear.Element + c.MustSetRandom() + + // add constant + f.AddConstantInPlace(&c) + + // check + var expectedCoeffs, one babybear.Element + one.SetOne() + expectedCoeffs.Add(&one, &c) + for i := range 20 { + if !f[i].Equal(&expectedCoeffs) { + t.Fatal("AddConstantInPlace failed") + } + } +} + +func TestPolynomialSubConstantInPlace(t *testing.T) { + + // build polynomial + f := make(Polynomial, 20) + for i := range 20 { + f[i].SetOne() + } + + // constant to sub + var c babybear.Element + c.MustSetRandom() + + // sub constant + f.SubConstantInPlace(&c) + + // check + var expectedCoeffs, one babybear.Element + one.SetOne() + expectedCoeffs.Sub(&one, &c) + for i := range 20 { + if !f[i].Equal(&expectedCoeffs) { + t.Fatal("SubConstantInPlace failed") + } + } +} + +func TestPolynomialScaleInPlace(t *testing.T) { + + // build polynomial + f := make(Polynomial, 20) + for i := range 20 { + f[i].SetOne() + } + + // constant to scale by + var c babybear.Element + c.MustSetRandom() + + // scale by constant + f.ScaleInPlace(&c) + + // check + for i := range 20 { + if !f[i].Equal(&c) { + t.Fatal("ScaleInPlace failed") + } + } + +} + +func TestPolynomialAdd(t *testing.T) { + + // build unbalanced polynomials + f1 := make(Polynomial, 20) + f1Backup := make(Polynomial, 20) + for i := range 20 { + f1[i].SetOne() + f1Backup[i].SetOne() + } + f2 := make(Polynomial, 10) + f2Backup := make(Polynomial, 10) + for i := range 10 { + f2[i].SetOne() + f2Backup[i].SetOne() + } + + // expected result + var one, two babybear.Element + one.SetOne() + two.Double(&one) + expectedSum := make(Polynomial, 20) + for i := range 10 { + expectedSum[i].Set(&two) + } + for i := 10; i < 20; i++ { + expectedSum[i].Set(&one) + } + + // caller is empty + var g Polynomial + g.Add(f1, f2) + if !g.Equal(expectedSum) { + t.Fatal("add polynomials fails") + } + if !f1.Equal(f1Backup) { + t.Fatal("side effect, f1 should not have been modified") + } + if !f2.Equal(f2Backup) { + t.Fatal("side effect, f2 should not have been modified") + } + + // all operands are distinct + _f1 := f1.Clone() + _f1.Add(f1, f2) + if !_f1.Equal(expectedSum) { + t.Fatal("add polynomials fails") + } + if !f1.Equal(f1Backup) { + t.Fatal("side effect, f1 should not have been modified") + } + if !f2.Equal(f2Backup) { + t.Fatal("side effect, f2 should not have been modified") + } + + // first operand = caller + _f1 = f1.Clone() + _f2 := f2.Clone() + _f1.Add(_f1, _f2) + if !_f1.Equal(expectedSum) { + t.Fatal("add polynomials fails") + } + if !_f2.Equal(f2Backup) { + t.Fatal("side effect, _f2 should not have been modified") + } + + // second operand = caller + _f1 = f1.Clone() + _f2 = f2.Clone() + _f1.Add(_f2, _f1) + if !_f1.Equal(expectedSum) { + t.Fatal("add polynomials fails") + } + if !_f2.Equal(f2Backup) { + t.Fatal("side effect, _f2 should not have been modified") + } +} + +func TestPolynomialText(t *testing.T) { + var one, negTwo babybear.Element + one.SetOne() + negTwo.SetInt64(-2) + + p := Polynomial{one, negTwo, one} + + assert.Equal(t, "X² - 2X + 1", p.Text(10)) +} + +func TestPrecomputeLagrange(t *testing.T) { + + testForDomainSize := func(domainSize uint8) bool { + polys := computeLagrangeBasis(domainSize) + + for l := range domainSize { + for i := range domainSize { + var I babybear.Element + I.SetUint64(uint64(i)) + y := polys[l].Eval(&I) + + if i == l && !y.IsOne() || i != l && !y.IsZero() { + t.Errorf("domainSize = %d: p_%d(%d) = %s", domainSize, l, i, y.Text(10)) + return false + } + } + } + return true + } + + t.Parallel() + parameters := gopter.DefaultTestParameters() + + const maxLagrangeDomainSize = 12 + + parameters.MinSuccessfulTests = maxLagrangeDomainSize + + properties := gopter.NewProperties(parameters) + + properties.Property("l'th lagrange polynomials must evaluate to 1 on l and 0 on other values in the domain", prop.ForAll( + testForDomainSize, + gen.UInt8Range(2, maxLagrangeDomainSize), + )) + + properties.TestingRun(t, gopter.ConsoleReporter(false)) +} + +func TestLagrangeCache(t *testing.T) { + for _, i := range []int{5, 2, 8, 4, 6, 3, 0} { + b := getLagrangeBasis(uint8(i)) + assert.Equal(t, b, getLagrangeBasis(uint8(i))) // second call must yield the same result + } +} diff --git a/field/babybear/polynomial/pool.go b/field/babybear/polynomial/pool.go new file mode 100644 index 0000000000..6a03e14325 --- /dev/null +++ b/field/babybear/polynomial/pool.go @@ -0,0 +1,191 @@ +// Copyright 2020-2026 Consensys Software Inc. +// Licensed under the Apache License, Version 2.0. See the LICENSE file for details. + +// Code generated by consensys/gnark-crypto DO NOT EDIT + +package polynomial + +import ( + "encoding/json" + "fmt" + "runtime" + "sort" + "sync" + "unsafe" + + "github.com/consensys/gnark-crypto/field/babybear" +) + +// Memory management for polynomials +// WARNING: This is not thread safe TODO: Make sure that is not a problem +// TODO: There is a lot of "unsafe" memory management here and needs to be vetted thoroughly + +type sizedPool struct { + maxN int + pool sync.Pool + stats poolStats +} + +type inUseData struct { + allocatedFor []uintptr + pool *sizedPool +} + +type Pool struct { + //lock sync.Mutex + inUse sync.Map + subPools []sizedPool +} + +func (p *sizedPool) get(n int) *babybear.Element { + p.stats.make(n) + return p.pool.Get().(*babybear.Element) +} + +func (p *sizedPool) put(ptr *babybear.Element) { + p.stats.dump() + p.pool.Put(ptr) +} + +func NewPool(maxN ...int) (pool Pool) { + + sort.Ints(maxN) + pool = Pool{ + subPools: make([]sizedPool, len(maxN)), + } + + for i := range pool.subPools { + subPool := &pool.subPools[i] + subPool.maxN = maxN[i] + subPool.pool = sync.Pool{ + New: func() any { + subPool.stats.Allocated++ + return getDataPointer(make([]babybear.Element, 0, subPool.maxN)) + }, + } + } + return +} + +func (p *Pool) findCorrespondingPool(n int) *sizedPool { + poolI := 0 + for poolI < len(p.subPools) && n > p.subPools[poolI].maxN { + poolI++ + } + return &p.subPools[poolI] // out of bounds error here would mean that n is too large +} + +func (p *Pool) Make(n int) []babybear.Element { + pool := p.findCorrespondingPool(n) + ptr := pool.get(n) + p.addInUse(ptr, pool) + return unsafe.Slice(ptr, n) +} + +// Dump dumps a set of polynomials into the pool +func (p *Pool) Dump(slices ...[]babybear.Element) { + for _, slice := range slices { + ptr := getDataPointer(slice) + if metadata, ok := p.inUse.Load(ptr); ok { + p.inUse.Delete(ptr) + metadata.(inUseData).pool.put(ptr) + } else { + panic("attempting to dump a slice not created by the pool") + } + } +} + +func (p *Pool) addInUse(ptr *babybear.Element, pool *sizedPool) { + pcs := make([]uintptr, 2) + n := runtime.Callers(3, pcs) + + if prevPcs, ok := p.inUse.Load(ptr); ok { // TODO: remove if unnecessary for security + panic(fmt.Errorf("re-allocated non-dumped slice, previously allocated at %v", runtime.CallersFrames(prevPcs.(inUseData).allocatedFor))) + } + p.inUse.Store(ptr, inUseData{ + allocatedFor: pcs[:n], + pool: pool, + }) +} + +func printFrame(frame runtime.Frame) { + fmt.Printf("\t%s line %d, function %s\n", frame.File, frame.Line, frame.Function) +} + +func (p *Pool) printInUse() { + fmt.Println("slices never dumped allocated at:") + p.inUse.Range(func(_, pcs any) bool { + fmt.Println("-------------------------") + + var frame runtime.Frame + frames := runtime.CallersFrames(pcs.(inUseData).allocatedFor) + more := true + for more { + frame, more = frames.Next() + printFrame(frame) + } + return true + }) +} + +type poolStats struct { + Used int + Allocated int + ReuseRate float64 + InUse int + GreatestNUsed int + SmallestNUsed int +} + +type poolsStats struct { + SubPools []poolStats + InUse int +} + +func (s *poolStats) make(n int) { + s.Used++ + s.InUse++ + if n > s.GreatestNUsed { + s.GreatestNUsed = n + } + if s.SmallestNUsed == 0 || s.SmallestNUsed > n { + s.SmallestNUsed = n + } +} + +func (s *poolStats) dump() { + s.InUse-- +} + +func (s *poolStats) finalize() { + s.ReuseRate = float64(s.Used) / float64(s.Allocated) +} + +func getDataPointer(slice []babybear.Element) *babybear.Element { + return (*babybear.Element)(unsafe.SliceData(slice)) +} + +func (p *Pool) PrintPoolStats() { + InUse := 0 + subStats := make([]poolStats, len(p.subPools)) + for i := range p.subPools { + subPool := &p.subPools[i] + subPool.stats.finalize() + subStats[i] = subPool.stats + InUse += subPool.stats.InUse + } + + stats := poolsStats{ + SubPools: subStats, + InUse: InUse, + } + serialized, _ := json.MarshalIndent(stats, "", " ") + fmt.Println(string(serialized)) + p.printInUse() +} + +func (p *Pool) Clone(slice []babybear.Element) []babybear.Element { + res := p.Make(len(slice)) + copy(res, slice) + return res +} diff --git a/field/goldilocks/extensions/e2.go b/field/goldilocks/extensions/e2.go index 35345a29c6..fbffc7d3ff 100644 --- a/field/goldilocks/extensions/e2.go +++ b/field/goldilocks/extensions/e2.go @@ -89,6 +89,30 @@ func (z *E2) SetUint64(v uint64) *E2 { return z } +// AddElement sets z to x + y, where y is an element of the base field embedded in E2, and returns z +func (z *E2) AddElement(x *E2, y *fr.Element) *E2 { + yc := *y + z.A0.Add(&x.A0, &yc) + z.A1 = x.A1 + return z +} + +// SubElement sets z to x - y, where y is an element of the base field embedded in E2, and returns z +func (z *E2) SubElement(x *E2, y *fr.Element) *E2 { + yc := *y + z.A0.Sub(&x.A0, &yc) + z.A1 = x.A1 + return z +} + +// SetElement sets z to x, an element of the base field embedded in E2, and returns z +func (z *E2) SetElement(x *fr.Element) *E2 { + v := *x + *z = E2{} + z.A0 = v + return z +} + // SetBigInt sets z to v (embedded via the unique ring homomorphism ℤ→𝔽, reduced modulo the base field order) and returns z func (z *E2) SetBigInt(v *big.Int) *E2 { *z = E2{} diff --git a/field/goldilocks/extensions/e2_test.go b/field/goldilocks/extensions/e2_test.go index cda8d7adb5..730bc30a1c 100644 --- a/field/goldilocks/extensions/e2_test.go +++ b/field/goldilocks/extensions/e2_test.go @@ -615,3 +615,34 @@ func TestE2SetBytesCanonicalRejectsNonCanonical(t *testing.T) { require.True(t, y.Equal(&x), "E2 must be unchanged on error") } } + +func TestE2ElementOps(t *testing.T) { + for range 100 { + var x, lifted, got, want E2 + var e fr.Element + x.MustSetRandom() + e.MustSetRandom() + + var z E2 + z.SetElement(&e) + require.True(t, z.A0.Equal(&e)) + require.True(t, z.A1.IsZero()) + lifted = z + + got.AddElement(&x, &e) + want.Add(&x, &lifted) + require.True(t, got.Equal(&want), "AddElement") + + got.SubElement(&x, &e) + want.Sub(&x, &lifted) + require.True(t, got.Equal(&want), "SubElement") + + // aliasing of the receiver with the first operand + got.Set(&x) + got.AddElement(&got, &e) + want.Add(&x, &lifted) + require.True(t, got.Equal(&want), "AddElement alias") + got.SubElement(&got, &e) + require.True(t, got.Equal(&x), "SubElement alias") + } +} diff --git a/field/goldilocks/polynomial/doc.go b/field/goldilocks/polynomial/doc.go new file mode 100644 index 0000000000..aa346f3ea3 --- /dev/null +++ b/field/goldilocks/polynomial/doc.go @@ -0,0 +1,7 @@ +// Copyright 2020-2026 Consensys Software Inc. +// Licensed under the Apache License, Version 2.0. See the LICENSE file for details. + +// Code generated by consensys/gnark-crypto DO NOT EDIT + +// Package polynomial provides polynomial methods and commitment schemes. +package polynomial diff --git a/field/goldilocks/polynomial/multilin.go b/field/goldilocks/polynomial/multilin.go new file mode 100644 index 0000000000..bfcc0f96f6 --- /dev/null +++ b/field/goldilocks/polynomial/multilin.go @@ -0,0 +1,179 @@ +// Copyright 2020-2026 Consensys Software Inc. +// Licensed under the Apache License, Version 2.0. See the LICENSE file for details. + +// Code generated by consensys/gnark-crypto DO NOT EDIT + +package polynomial + +import ( + "math/bits" + + "github.com/consensys/gnark-crypto/field/goldilocks" + "github.com/consensys/gnark-crypto/utils" +) + +// MultiLin tracks the values of a (dense i.e. not sparse) multilinear polynomial +// The variables are X₁ through Xₙ where n = log(len(.)) +// .[∑ᵢ 2ⁱ⁻¹ bₙ₋ᵢ] = the polynomial evaluated at (b₁, b₂, ..., bₙ) +// It is understood that any hypercube evaluation can be extrapolated to a multilinear polynomial +type MultiLin []goldilocks.Element + +// Fold is partial evaluation function k[X₁, X₂, ..., Xₙ] → k[X₂, ..., Xₙ] by setting X₁=r +func (m *MultiLin) Fold(r goldilocks.Element) { + mid := len(*m) / 2 + + bottom, top := (*m)[:mid], (*m)[mid:] + + var t goldilocks.Element // no need to update the top part + + // updating bookkeeping table + // knowing that the polynomial f ∈ (k[X₂, ..., Xₙ])[X₁] is linear, we would get f(r) = f(0) + r(f(1) - f(0)) + // the following loop computes the evaluations of f(r) accordingly: + // f(r, b₂, ..., bₙ) = f(0, b₂, ..., bₙ) + r(f(1, b₂, ..., bₙ) - f(0, b₂, ..., bₙ)) + for i := range mid { + // table[i] ← table[i] + r (table[i + mid] - table[i]) + t.Sub(&top[i], &bottom[i]) + t.Mul(&t, &r) + bottom[i].Add(&bottom[i], &t) + } + + *m = (*m)[:mid] +} + +func (m *MultiLin) FoldParallel(r goldilocks.Element) utils.Task { + mid := len(*m) / 2 + bottom, top := (*m)[:mid], (*m)[mid:] + + *m = bottom + + return func(start, end int) { + var t goldilocks.Element // no need to update the top part + for i := start; i < end; i++ { + // table[i] ← table[i] + r (table[i + mid] - table[i]) + t.Sub(&top[i], &bottom[i]) + t.Mul(&t, &r) + bottom[i].Add(&bottom[i], &t) + } + } +} + +func (m MultiLin) Sum() goldilocks.Element { + s := m[0] + for i := 1; i < len(m); i++ { + s.Add(&s, &m[i]) + } + return s +} + +func _clone(m MultiLin, p *Pool) MultiLin { + if p == nil { + return m.Clone() + } else { + return p.Clone(m) + } +} + +func _dump(m MultiLin, p *Pool) { + if p != nil { + p.Dump(m) + } +} + +// Evaluate extrapolate the value of the multilinear polynomial corresponding to m +// on the given coordinates +func (m MultiLin) Evaluate(coordinates []goldilocks.Element, p *Pool) goldilocks.Element { + // Folding is a mutating operation + bkCopy := _clone(m, p) + + // Evaluate step by step through repeated folding (i.e. evaluation at the first remaining variable) + for _, r := range coordinates { + bkCopy.Fold(r) + } + + result := bkCopy[0] + + _dump(bkCopy, p) + return result +} + +// Clone creates a deep copy of a bookkeeping table. +// Both multilinear interpolation and sumcheck require folding an underlying +// array, but folding changes the array. To do both one requires a deep copy +// of the bookkeeping table. +func (m MultiLin) Clone() MultiLin { + res := make(MultiLin, len(m)) + copy(res, m) + return res +} + +// Add two bookKeepingTables +func (m *MultiLin) Add(left, right MultiLin) { + size := len(left) + // Check that left and right have the same size + if len(right) != size || len(*m) != size { + panic("left, right and destination must have the right size") + } + + // Add elementwise + for i := range size { + (*m)[i].Add(&left[i], &right[i]) + } +} + +// EvalEq computes Eq(q₁, ... , qₙ, h₁, ... , hₙ) = Π₁ⁿ Eq(qᵢ, hᵢ) +// where Eq(x,y) = xy + (1-x)(1-y) = 1 - x - y + xy + xy interpolates +// +// _________________ +// | | | +// | 0 | 1 | +// |_______|_______| +// y | | | +// | 1 | 0 | +// |_______|_______| +// +// x +// +// In other words the polynomial evaluated here is the multilinear extrapolation of +// one that evaluates to q' == h' for vectors q', h' of binary values +func EvalEq(q, h []goldilocks.Element) goldilocks.Element { + var res, nxt, one, sum goldilocks.Element + one.SetOne() + for i := range len(q) { + nxt.Mul(&q[i], &h[i]) // nxt <- qᵢ * hᵢ + nxt.Double(&nxt) // nxt <- 2 * qᵢ * hᵢ + nxt.Add(&nxt, &one) // nxt <- 1 + 2 * qᵢ * hᵢ + sum.Add(&q[i], &h[i]) // sum <- qᵢ + hᵢ TODO: Why not subtract one by one from nxt? More parallel? + + if i == 0 { + res.Sub(&nxt, &sum) // nxt <- 1 + 2 * qᵢ * hᵢ - qᵢ - hᵢ + } else { + nxt.Sub(&nxt, &sum) // nxt <- 1 + 2 * qᵢ * hᵢ - qᵢ - hᵢ + res.Mul(&res, &nxt) // res <- res * nxt + } + } + return res +} + +// Eq sets m to the representation of the polynomial Eq(q₁, ..., qₙ, *, ..., *) × m[0] +func (m *MultiLin) Eq(q []goldilocks.Element) { + n := len(q) + + if len(*m) != 1<= 0; i-- { + res.Mul(&res, v) + res.Add(&res, &(*p)[i]) + } + + return res +} + +// Clone returns a copy of the polynomial +func (p *Polynomial) Clone() Polynomial { + _p := make(Polynomial, len(*p)) + copy(_p, *p) + return _p +} + +// Set to another polynomial +func (p *Polynomial) Set(p1 Polynomial) { + if len(*p) != len(p1) { + *p = p1.Clone() + return + } + + for i := range len(p1) { + (*p)[i].Set(&p1[i]) + } +} + +// AddConstantInPlace adds a constant to the polynomial, modifying p +func (p *Polynomial) AddConstantInPlace(c *goldilocks.Element) { + for i := range len(*p) { + (*p)[i].Add(&(*p)[i], c) + } +} + +// SubConstantInPlace subs a constant to the polynomial, modifying p +func (p *Polynomial) SubConstantInPlace(c *goldilocks.Element) { + for i := range len(*p) { + (*p)[i].Sub(&(*p)[i], c) + } +} + +// ScaleInPlace multiplies p by v, modifying p +func (p *Polynomial) ScaleInPlace(c *goldilocks.Element) { + for i := range len(*p) { + (*p)[i].Mul(&(*p)[i], c) + } +} + +// Scale multiplies p0 by v, storing the result in p +func (p *Polynomial) Scale(c *goldilocks.Element, p0 Polynomial) { + if len(*p) != len(p0) { + *p = make(Polynomial, len(p0)) + } + for i := range len(p0) { + (*p)[i].Mul(c, &p0[i]) + } +} + +// Add adds p1 to p2 +// This function allocates a new slice unless p == p1 or p == p2 +func (p *Polynomial) Add(p1, p2 Polynomial) *Polynomial { + + bigger := p1 + smaller := p2 + if len(bigger) < len(smaller) { + bigger, smaller = smaller, bigger + } + + if len(*p) == len(bigger) && (&(*p)[0] == &bigger[0]) { + for i := range len(smaller) { + (*p)[i].Add(&(*p)[i], &smaller[i]) + } + return p + } + + if len(*p) == len(smaller) && (&(*p)[0] == &smaller[0]) { + for i := range len(smaller) { + (*p)[i].Add(&(*p)[i], &bigger[i]) + } + *p = append(*p, bigger[len(smaller):]...) + return p + } + + res := make(Polynomial, len(bigger)) + copy(res, bigger) + for i := range len(smaller) { + res[i].Add(&res[i], &smaller[i]) + } + *p = res + return p +} + +// Sub subtracts p2 from p1 +// TODO make interface more consistent with Add +func (p *Polynomial) Sub(p1, p2 Polynomial) *Polynomial { + if len(p1) != len(p2) || len(p2) != len(*p) { + return nil + } + for i := range len(*p) { + (*p)[i].Sub(&p1[i], &p2[i]) + } + return p +} + +// Equal checks equality between two polynomials +func (p *Polynomial) Equal(p1 Polynomial) bool { + if (*p == nil) != (p1 == nil) { + return false + } + + if len(*p) != len(p1) { + return false + } + + for i := range p1 { + if !(*p)[i].Equal(&p1[i]) { + return false + } + } + + return true +} + +func (p Polynomial) SetZero() { + for i := range len(p) { + p[i].SetZero() + } +} + +func (p Polynomial) Text(base int) string { + + var builder strings.Builder + + first := true + for d := len(p) - 1; d >= 0; d-- { + if p[d].IsZero() { + continue + } + + pD := p[d] + pDText := pD.Text(base) + + initialLen := builder.Len() + + if pDText[0] == '-' { + pDText = pDText[1:] + if first { + builder.WriteString("-") + } else { + builder.WriteString(" - ") + } + } else if !first { + builder.WriteString(" + ") + } + + first = false + + if !pD.IsOne() || d == 0 { + builder.WriteString(pDText) + } + + if builder.Len()-initialLen > 10 { + builder.WriteString("×") + } + + if d != 0 { + builder.WriteString("X") + } + if d > 1 { + builder.WriteString( + utils.ToSuperscript(strconv.Itoa(d)), + ) + } + + } + + if first { + return "0" + } + + return builder.String() +} + +// InterpolateOnRange maps vector v to polynomial f +// such that f(i) = v[i] for 0 ≤ i < len(v). +// len(f) = len(v) and deg(f) ≤ len(v) - 1 +func InterpolateOnRange(v []goldilocks.Element) Polynomial { + nEvals := uint8(len(v)) + if int(nEvals) != len(v) { + panic("interpolation method too inefficient for nEvals > 255") + } + lagrange := getLagrangeBasis(nEvals) + + var res Polynomial + res.Scale(&v[0], lagrange[0]) + + temp := make(Polynomial, nEvals) + + for i := uint8(1); i < nEvals; i++ { + temp.Scale(&v[i], lagrange[i]) + res.Add(res, temp) + } + + return res +} + +// lagrange bases used by InterpolateOnRange +var lagrangeBasis sync.Map + +func getLagrangeBasis(domainSize uint8) []Polynomial { + if res, ok := lagrangeBasis.Load(domainSize); ok { + return res.([]Polynomial) + } + + // not found. compute + var res []Polynomial + if domainSize >= 2 { + res = computeLagrangeBasis(domainSize) + } else if domainSize == 1 { + res = []Polynomial{make(Polynomial, 1)} + res[0][0].SetOne() + } + lagrangeBasis.Store(domainSize, res) + + return res +} + +// computeLagrangeBasis precomputes in explicit coefficient form for each 0 ≤ l < domainSize the polynomial +// pₗ := X (X-1) ... (X-l-1) (X-l+1) ... (X - domainSize + 1) / ( l (l-1) ... 2 (-1) ... (l - domainSize +1) ) +// Note that pₗ(l) = 1 and pₗ(n) = 0 if 0 ≤ l < domainSize, n ≠ l +func computeLagrangeBasis(domainSize uint8) []Polynomial { + + constTerms := make([]goldilocks.Element, domainSize) + for i := range domainSize { + constTerms[i].SetInt64(-int64(i)) + } + + res := make([]Polynomial, domainSize) + multScratch := make(Polynomial, domainSize-1) + + // compute pₗ + for l := range domainSize { + + // TODO @Tabaie Optimize this with some trees? O(log(domainSize)) polynomial mults instead of O(domainSize)? Then again it would be fewer big poly mults vs many small poly mults + d := uint8(0) //d is the current degree of res + for i := range domainSize { + if i == l { + continue + } + if d == 0 { + res[l] = make(Polynomial, domainSize) + res[l][domainSize-2] = constTerms[i] + res[l][domainSize-1].SetOne() + } else { + current := res[l][domainSize-d-2:] + timesConst := multScratch[domainSize-d-2:] + + timesConst.Scale(&constTerms[i], current[1:]) //TODO: Directly double and add since constTerms are tiny? (even less than 4 bits) + nonLeading := current[0 : d+1] + + nonLeading.Add(nonLeading, timesConst) + + } + d++ + } + + } + + // We have pₗ(i≠l)=0. Now scale so that pₗ(l)=1 + // Replace the constTerms with norms + for l := range domainSize { + constTerms[l].Neg(&constTerms[l]) + constTerms[l] = res[l].Eval(&constTerms[l]) + } + constTerms = goldilocks.BatchInvert(constTerms) + for l := range domainSize { + res[l].ScaleInPlace(&constTerms[l]) + } + + return res +} diff --git a/field/goldilocks/polynomial/polynomial_test.go b/field/goldilocks/polynomial/polynomial_test.go new file mode 100644 index 0000000000..3a36ae3a53 --- /dev/null +++ b/field/goldilocks/polynomial/polynomial_test.go @@ -0,0 +1,255 @@ +// Copyright 2020-2026 Consensys Software Inc. +// Licensed under the Apache License, Version 2.0. See the LICENSE file for details. + +// Code generated by consensys/gnark-crypto DO NOT EDIT + +package polynomial + +import ( + "math/big" + "testing" + + "github.com/consensys/gnark-crypto/field/goldilocks" + "github.com/leanovate/gopter" + "github.com/leanovate/gopter/gen" + "github.com/leanovate/gopter/prop" + "github.com/stretchr/testify/assert" +) + +func TestPolynomialEval(t *testing.T) { + + // build polynomial + f := make(Polynomial, 20) + for i := range 20 { + f[i].SetOne() + } + + // random value + var point goldilocks.Element + point.MustSetRandom() + + // compute manually f(val) + var expectedEval, one, den goldilocks.Element + var expo big.Int + one.SetOne() + expo.SetUint64(20) + expectedEval.Exp(point, &expo). + Sub(&expectedEval, &one) + den.Sub(&point, &one) + expectedEval.Div(&expectedEval, &den) + + // compute purported evaluation + purportedEval := f.Eval(&point) + + // check + if !purportedEval.Equal(&expectedEval) { + t.Fatal("polynomial evaluation failed") + } +} + +func TestPolynomialAddConstantInPlace(t *testing.T) { + + // build polynomial + f := make(Polynomial, 20) + for i := range 20 { + f[i].SetOne() + } + + // constant to add + var c goldilocks.Element + c.MustSetRandom() + + // add constant + f.AddConstantInPlace(&c) + + // check + var expectedCoeffs, one goldilocks.Element + one.SetOne() + expectedCoeffs.Add(&one, &c) + for i := range 20 { + if !f[i].Equal(&expectedCoeffs) { + t.Fatal("AddConstantInPlace failed") + } + } +} + +func TestPolynomialSubConstantInPlace(t *testing.T) { + + // build polynomial + f := make(Polynomial, 20) + for i := range 20 { + f[i].SetOne() + } + + // constant to sub + var c goldilocks.Element + c.MustSetRandom() + + // sub constant + f.SubConstantInPlace(&c) + + // check + var expectedCoeffs, one goldilocks.Element + one.SetOne() + expectedCoeffs.Sub(&one, &c) + for i := range 20 { + if !f[i].Equal(&expectedCoeffs) { + t.Fatal("SubConstantInPlace failed") + } + } +} + +func TestPolynomialScaleInPlace(t *testing.T) { + + // build polynomial + f := make(Polynomial, 20) + for i := range 20 { + f[i].SetOne() + } + + // constant to scale by + var c goldilocks.Element + c.MustSetRandom() + + // scale by constant + f.ScaleInPlace(&c) + + // check + for i := range 20 { + if !f[i].Equal(&c) { + t.Fatal("ScaleInPlace failed") + } + } + +} + +func TestPolynomialAdd(t *testing.T) { + + // build unbalanced polynomials + f1 := make(Polynomial, 20) + f1Backup := make(Polynomial, 20) + for i := range 20 { + f1[i].SetOne() + f1Backup[i].SetOne() + } + f2 := make(Polynomial, 10) + f2Backup := make(Polynomial, 10) + for i := range 10 { + f2[i].SetOne() + f2Backup[i].SetOne() + } + + // expected result + var one, two goldilocks.Element + one.SetOne() + two.Double(&one) + expectedSum := make(Polynomial, 20) + for i := range 10 { + expectedSum[i].Set(&two) + } + for i := 10; i < 20; i++ { + expectedSum[i].Set(&one) + } + + // caller is empty + var g Polynomial + g.Add(f1, f2) + if !g.Equal(expectedSum) { + t.Fatal("add polynomials fails") + } + if !f1.Equal(f1Backup) { + t.Fatal("side effect, f1 should not have been modified") + } + if !f2.Equal(f2Backup) { + t.Fatal("side effect, f2 should not have been modified") + } + + // all operands are distinct + _f1 := f1.Clone() + _f1.Add(f1, f2) + if !_f1.Equal(expectedSum) { + t.Fatal("add polynomials fails") + } + if !f1.Equal(f1Backup) { + t.Fatal("side effect, f1 should not have been modified") + } + if !f2.Equal(f2Backup) { + t.Fatal("side effect, f2 should not have been modified") + } + + // first operand = caller + _f1 = f1.Clone() + _f2 := f2.Clone() + _f1.Add(_f1, _f2) + if !_f1.Equal(expectedSum) { + t.Fatal("add polynomials fails") + } + if !_f2.Equal(f2Backup) { + t.Fatal("side effect, _f2 should not have been modified") + } + + // second operand = caller + _f1 = f1.Clone() + _f2 = f2.Clone() + _f1.Add(_f2, _f1) + if !_f1.Equal(expectedSum) { + t.Fatal("add polynomials fails") + } + if !_f2.Equal(f2Backup) { + t.Fatal("side effect, _f2 should not have been modified") + } +} + +func TestPolynomialText(t *testing.T) { + var one, negTwo goldilocks.Element + one.SetOne() + negTwo.SetInt64(-2) + + p := Polynomial{one, negTwo, one} + + assert.Equal(t, "X² - 2X + 1", p.Text(10)) +} + +func TestPrecomputeLagrange(t *testing.T) { + + testForDomainSize := func(domainSize uint8) bool { + polys := computeLagrangeBasis(domainSize) + + for l := range domainSize { + for i := range domainSize { + var I goldilocks.Element + I.SetUint64(uint64(i)) + y := polys[l].Eval(&I) + + if i == l && !y.IsOne() || i != l && !y.IsZero() { + t.Errorf("domainSize = %d: p_%d(%d) = %s", domainSize, l, i, y.Text(10)) + return false + } + } + } + return true + } + + t.Parallel() + parameters := gopter.DefaultTestParameters() + + const maxLagrangeDomainSize = 12 + + parameters.MinSuccessfulTests = maxLagrangeDomainSize + + properties := gopter.NewProperties(parameters) + + properties.Property("l'th lagrange polynomials must evaluate to 1 on l and 0 on other values in the domain", prop.ForAll( + testForDomainSize, + gen.UInt8Range(2, maxLagrangeDomainSize), + )) + + properties.TestingRun(t, gopter.ConsoleReporter(false)) +} + +func TestLagrangeCache(t *testing.T) { + for _, i := range []int{5, 2, 8, 4, 6, 3, 0} { + b := getLagrangeBasis(uint8(i)) + assert.Equal(t, b, getLagrangeBasis(uint8(i))) // second call must yield the same result + } +} diff --git a/field/goldilocks/polynomial/pool.go b/field/goldilocks/polynomial/pool.go new file mode 100644 index 0000000000..7ee7ee2213 --- /dev/null +++ b/field/goldilocks/polynomial/pool.go @@ -0,0 +1,191 @@ +// Copyright 2020-2026 Consensys Software Inc. +// Licensed under the Apache License, Version 2.0. See the LICENSE file for details. + +// Code generated by consensys/gnark-crypto DO NOT EDIT + +package polynomial + +import ( + "encoding/json" + "fmt" + "runtime" + "sort" + "sync" + "unsafe" + + "github.com/consensys/gnark-crypto/field/goldilocks" +) + +// Memory management for polynomials +// WARNING: This is not thread safe TODO: Make sure that is not a problem +// TODO: There is a lot of "unsafe" memory management here and needs to be vetted thoroughly + +type sizedPool struct { + maxN int + pool sync.Pool + stats poolStats +} + +type inUseData struct { + allocatedFor []uintptr + pool *sizedPool +} + +type Pool struct { + //lock sync.Mutex + inUse sync.Map + subPools []sizedPool +} + +func (p *sizedPool) get(n int) *goldilocks.Element { + p.stats.make(n) + return p.pool.Get().(*goldilocks.Element) +} + +func (p *sizedPool) put(ptr *goldilocks.Element) { + p.stats.dump() + p.pool.Put(ptr) +} + +func NewPool(maxN ...int) (pool Pool) { + + sort.Ints(maxN) + pool = Pool{ + subPools: make([]sizedPool, len(maxN)), + } + + for i := range pool.subPools { + subPool := &pool.subPools[i] + subPool.maxN = maxN[i] + subPool.pool = sync.Pool{ + New: func() any { + subPool.stats.Allocated++ + return getDataPointer(make([]goldilocks.Element, 0, subPool.maxN)) + }, + } + } + return +} + +func (p *Pool) findCorrespondingPool(n int) *sizedPool { + poolI := 0 + for poolI < len(p.subPools) && n > p.subPools[poolI].maxN { + poolI++ + } + return &p.subPools[poolI] // out of bounds error here would mean that n is too large +} + +func (p *Pool) Make(n int) []goldilocks.Element { + pool := p.findCorrespondingPool(n) + ptr := pool.get(n) + p.addInUse(ptr, pool) + return unsafe.Slice(ptr, n) +} + +// Dump dumps a set of polynomials into the pool +func (p *Pool) Dump(slices ...[]goldilocks.Element) { + for _, slice := range slices { + ptr := getDataPointer(slice) + if metadata, ok := p.inUse.Load(ptr); ok { + p.inUse.Delete(ptr) + metadata.(inUseData).pool.put(ptr) + } else { + panic("attempting to dump a slice not created by the pool") + } + } +} + +func (p *Pool) addInUse(ptr *goldilocks.Element, pool *sizedPool) { + pcs := make([]uintptr, 2) + n := runtime.Callers(3, pcs) + + if prevPcs, ok := p.inUse.Load(ptr); ok { // TODO: remove if unnecessary for security + panic(fmt.Errorf("re-allocated non-dumped slice, previously allocated at %v", runtime.CallersFrames(prevPcs.(inUseData).allocatedFor))) + } + p.inUse.Store(ptr, inUseData{ + allocatedFor: pcs[:n], + pool: pool, + }) +} + +func printFrame(frame runtime.Frame) { + fmt.Printf("\t%s line %d, function %s\n", frame.File, frame.Line, frame.Function) +} + +func (p *Pool) printInUse() { + fmt.Println("slices never dumped allocated at:") + p.inUse.Range(func(_, pcs any) bool { + fmt.Println("-------------------------") + + var frame runtime.Frame + frames := runtime.CallersFrames(pcs.(inUseData).allocatedFor) + more := true + for more { + frame, more = frames.Next() + printFrame(frame) + } + return true + }) +} + +type poolStats struct { + Used int + Allocated int + ReuseRate float64 + InUse int + GreatestNUsed int + SmallestNUsed int +} + +type poolsStats struct { + SubPools []poolStats + InUse int +} + +func (s *poolStats) make(n int) { + s.Used++ + s.InUse++ + if n > s.GreatestNUsed { + s.GreatestNUsed = n + } + if s.SmallestNUsed == 0 || s.SmallestNUsed > n { + s.SmallestNUsed = n + } +} + +func (s *poolStats) dump() { + s.InUse-- +} + +func (s *poolStats) finalize() { + s.ReuseRate = float64(s.Used) / float64(s.Allocated) +} + +func getDataPointer(slice []goldilocks.Element) *goldilocks.Element { + return (*goldilocks.Element)(unsafe.SliceData(slice)) +} + +func (p *Pool) PrintPoolStats() { + InUse := 0 + subStats := make([]poolStats, len(p.subPools)) + for i := range p.subPools { + subPool := &p.subPools[i] + subPool.stats.finalize() + subStats[i] = subPool.stats + InUse += subPool.stats.InUse + } + + stats := poolsStats{ + SubPools: subStats, + InUse: InUse, + } + serialized, _ := json.MarshalIndent(stats, "", " ") + fmt.Println(string(serialized)) + p.printInUse() +} + +func (p *Pool) Clone(slice []goldilocks.Element) []goldilocks.Element { + res := p.Make(len(slice)) + copy(res, slice) + return res +} diff --git a/field/koalabear/extensions/e2.go b/field/koalabear/extensions/e2.go index db49bafcf7..964cfbccee 100644 --- a/field/koalabear/extensions/e2.go +++ b/field/koalabear/extensions/e2.go @@ -89,6 +89,30 @@ func (z *E2) SetUint64(v uint64) *E2 { return z } +// AddElement sets z to x + y, where y is an element of the base field embedded in E2, and returns z +func (z *E2) AddElement(x *E2, y *fr.Element) *E2 { + yc := *y + z.A0.Add(&x.A0, &yc) + z.A1 = x.A1 + return z +} + +// SubElement sets z to x - y, where y is an element of the base field embedded in E2, and returns z +func (z *E2) SubElement(x *E2, y *fr.Element) *E2 { + yc := *y + z.A0.Sub(&x.A0, &yc) + z.A1 = x.A1 + return z +} + +// SetElement sets z to x, an element of the base field embedded in E2, and returns z +func (z *E2) SetElement(x *fr.Element) *E2 { + v := *x + *z = E2{} + z.A0 = v + return z +} + // SetBigInt sets z to v (embedded via the unique ring homomorphism ℤ→𝔽, reduced modulo the base field order) and returns z func (z *E2) SetBigInt(v *big.Int) *E2 { *z = E2{} diff --git a/field/koalabear/extensions/e2_test.go b/field/koalabear/extensions/e2_test.go index df462eab32..4dea016c51 100644 --- a/field/koalabear/extensions/e2_test.go +++ b/field/koalabear/extensions/e2_test.go @@ -632,3 +632,34 @@ func TestE2SetBytesCanonicalRejectsNonCanonical(t *testing.T) { require.True(t, y.Equal(&x), "E2 must be unchanged on error") } } + +func TestE2ElementOps(t *testing.T) { + for range 100 { + var x, lifted, got, want E2 + var e fr.Element + x.MustSetRandom() + e.MustSetRandom() + + var z E2 + z.SetElement(&e) + require.True(t, z.A0.Equal(&e)) + require.True(t, z.A1.IsZero()) + lifted = z + + got.AddElement(&x, &e) + want.Add(&x, &lifted) + require.True(t, got.Equal(&want), "AddElement") + + got.SubElement(&x, &e) + want.Sub(&x, &lifted) + require.True(t, got.Equal(&want), "SubElement") + + // aliasing of the receiver with the first operand + got.Set(&x) + got.AddElement(&got, &e) + want.Add(&x, &lifted) + require.True(t, got.Equal(&want), "AddElement alias") + got.SubElement(&got, &e) + require.True(t, got.Equal(&x), "SubElement alias") + } +} diff --git a/field/koalabear/extensions/e4.go b/field/koalabear/extensions/e4.go index 31af799017..2ea799059d 100644 --- a/field/koalabear/extensions/e4.go +++ b/field/koalabear/extensions/e4.go @@ -100,6 +100,34 @@ func (z *E4) SetUint64(v uint64) *E4 { return z } +// AddElement sets z to x + y, where y is an element of the base field embedded in E4, and returns z +func (z *E4) AddElement(x *E4, y *fr.Element) *E4 { + yc := *y + z.B0.A0.Add(&x.B0.A0, &yc) + z.B0.A1 = x.B0.A1 + z.B1.A0 = x.B1.A0 + z.B1.A1 = x.B1.A1 + return z +} + +// SubElement sets z to x - y, where y is an element of the base field embedded in E4, and returns z +func (z *E4) SubElement(x *E4, y *fr.Element) *E4 { + yc := *y + z.B0.A0.Sub(&x.B0.A0, &yc) + z.B0.A1 = x.B0.A1 + z.B1.A0 = x.B1.A0 + z.B1.A1 = x.B1.A1 + return z +} + +// SetElement sets z to x, an element of the base field embedded in E4, and returns z +func (z *E4) SetElement(x *fr.Element) *E4 { + v := *x + *z = E4{} + z.B0.A0 = v + return z +} + // SetBigInt sets z to v (embedded via the unique ring homomorphism ℤ→𝔽, reduced modulo the base field order) and returns z func (z *E4) SetBigInt(v *big.Int) *E4 { *z = E4{} diff --git a/field/koalabear/extensions/e4_test.go b/field/koalabear/extensions/e4_test.go index c55631d70f..4ceea99856 100644 --- a/field/koalabear/extensions/e4_test.go +++ b/field/koalabear/extensions/e4_test.go @@ -1206,3 +1206,36 @@ func TestE4SetBytesCanonicalRejectsNonCanonical(t *testing.T) { require.True(t, y.Equal(&x), "E4 must be unchanged on error") } } + +func TestE4ElementOps(t *testing.T) { + for range 100 { + var x, lifted, got, want E4 + var e fr.Element + x.MustSetRandom() + e.MustSetRandom() + + var z E4 + z.SetElement(&e) + require.True(t, z.B0.A0.Equal(&e)) + require.True(t, z.B0.A1.IsZero()) + require.True(t, z.B1.A0.IsZero()) + require.True(t, z.B1.A1.IsZero()) + lifted = z + + got.AddElement(&x, &e) + want.Add(&x, &lifted) + require.True(t, got.Equal(&want), "AddElement") + + got.SubElement(&x, &e) + want.Sub(&x, &lifted) + require.True(t, got.Equal(&want), "SubElement") + + // aliasing of the receiver with the first operand + got.Set(&x) + got.AddElement(&got, &e) + want.Add(&x, &lifted) + require.True(t, got.Equal(&want), "AddElement alias") + got.SubElement(&got, &e) + require.True(t, got.Equal(&x), "SubElement alias") + } +} diff --git a/field/koalabear/extensions/e6.go b/field/koalabear/extensions/e6.go index c7df7659d0..8d3054a7a9 100644 --- a/field/koalabear/extensions/e6.go +++ b/field/koalabear/extensions/e6.go @@ -91,6 +91,38 @@ func (z *E6) SetUint64(v uint64) *E6 { return z } +// AddElement sets z to x + y, where y is an element of the base field embedded in E6, and returns z +func (z *E6) AddElement(x *E6, y *fr.Element) *E6 { + yc := *y + z.B0.A0.Add(&x.B0.A0, &yc) + z.B0.A1 = x.B0.A1 + z.B1.A0 = x.B1.A0 + z.B1.A1 = x.B1.A1 + z.B2.A0 = x.B2.A0 + z.B2.A1 = x.B2.A1 + return z +} + +// SubElement sets z to x - y, where y is an element of the base field embedded in E6, and returns z +func (z *E6) SubElement(x *E6, y *fr.Element) *E6 { + yc := *y + z.B0.A0.Sub(&x.B0.A0, &yc) + z.B0.A1 = x.B0.A1 + z.B1.A0 = x.B1.A0 + z.B1.A1 = x.B1.A1 + z.B2.A0 = x.B2.A0 + z.B2.A1 = x.B2.A1 + return z +} + +// SetElement sets z to x, an element of the base field embedded in E6, and returns z +func (z *E6) SetElement(x *fr.Element) *E6 { + v := *x + *z = E6{} + z.B0.A0 = v + return z +} + // SetBigInt sets z to v (embedded via the unique ring homomorphism ℤ→𝔽, reduced modulo the base field order) and returns z func (z *E6) SetBigInt(v *big.Int) *E6 { *z = E6{} diff --git a/field/koalabear/extensions/e6_test.go b/field/koalabear/extensions/e6_test.go index 02187bcabc..b14ada11f1 100644 --- a/field/koalabear/extensions/e6_test.go +++ b/field/koalabear/extensions/e6_test.go @@ -409,3 +409,38 @@ func TestE6SetBytesCanonicalRejectsNonCanonical(t *testing.T) { require.True(t, y.Equal(&x), "E6 must be unchanged on error") } } + +func TestE6ElementOps(t *testing.T) { + for range 100 { + var x, lifted, got, want E6 + var e fr.Element + x.MustSetRandom() + e.MustSetRandom() + + var z E6 + z.SetElement(&e) + require.True(t, z.B0.A0.Equal(&e)) + require.True(t, z.B0.A1.IsZero()) + require.True(t, z.B1.A0.IsZero()) + require.True(t, z.B1.A1.IsZero()) + require.True(t, z.B2.A0.IsZero()) + require.True(t, z.B2.A1.IsZero()) + lifted = z + + got.AddElement(&x, &e) + want.Add(&x, &lifted) + require.True(t, got.Equal(&want), "AddElement") + + got.SubElement(&x, &e) + want.Sub(&x, &lifted) + require.True(t, got.Equal(&want), "SubElement") + + // aliasing of the receiver with the first operand + got.Set(&x) + got.AddElement(&got, &e) + want.Add(&x, &lifted) + require.True(t, got.Equal(&want), "AddElement alias") + got.SubElement(&got, &e) + require.True(t, got.Equal(&x), "SubElement alias") + } +} diff --git a/field/koalabear/extensions/polynomial/multilin_e6.go b/field/koalabear/extensions/polynomial/multilin_e6.go index 929478d873..e9a833eb9f 100644 --- a/field/koalabear/extensions/polynomial/multilin_e6.go +++ b/field/koalabear/extensions/polynomial/multilin_e6.go @@ -8,7 +8,9 @@ package polynomial import ( "math/bits" + fr "github.com/consensys/gnark-crypto/field/koalabear" "github.com/consensys/gnark-crypto/field/koalabear/extensions" + basepoly "github.com/consensys/gnark-crypto/field/koalabear/polynomial" "github.com/consensys/gnark-crypto/utils" ) @@ -57,6 +59,67 @@ func (m *MultiLinE6) FoldParallel(r extensions.E6) utils.Task { } } +// FoldFromBase sets m to the partial evaluation X₁=r of the multilinear polynomial whose +// hypercube evaluations are the base field elements b: +// +// m[i] = b[i] + r (b[i + len(b)/2] - b[i]) +// +// m's backing array is reused if it is large enough. +func (m *MultiLinE6) FoldFromBase(b basepoly.MultiLin, r *extensions.E6) { + mid := len(b) / 2 + m.resize(mid) + m.foldFromBase(b, r, 0, mid) +} + +// FoldFromBaseParallel is the parallel version of FoldFromBase. It sizes m, and returns +// a task that computes the entries m[start:end]. The task reads *r when it runs, so *r must +// not be modified until all tasks have completed. +func (m *MultiLinE6) FoldFromBaseParallel(b basepoly.MultiLin, r *extensions.E6) utils.Task { + mid := len(b) / 2 + m.resize(mid) + dst := *m + + return func(start, end int) { + dst.foldFromBase(b, r, start, end) + } +} + +// EvaluateBase evaluates, on the given coordinates, the multilinear polynomial whose hypercube +// evaluations are the base field elements b. It first folds b into m by the first coordinate +// (see FoldFromBase), and then folds m by the remaining coordinates. m serves as scratch space, +// is overwritten, and its backing array is reused if it is large enough. +// Unlike Evaluate, no copy of the table is made, and b is left untouched. +func (m *MultiLinE6) EvaluateBase(b basepoly.MultiLin, coordinates []extensions.E6) extensions.E6 { + if len(coordinates) == 0 { + m.resize(1) + return *(*m)[0].SetElement(&b[0]) + } + + m.FoldFromBase(b, &coordinates[0]) + for _, r := range coordinates[1:] { + m.Fold(r) + } + return (*m)[0] +} + +func (m *MultiLinE6) resize(n int) { + if cap(*m) >= n { + *m = (*m)[:n] + } else { + *m = make(MultiLinE6, n) + } +} + +func (m MultiLinE6) foldFromBase(b basepoly.MultiLin, r *extensions.E6, start, end int) { + mid := len(b) / 2 + var diff fr.Element + for i := start; i < end; i++ { + diff.Sub(&b[mid+i], &b[i]) + m[i].MulByElement(r, &diff) + m[i].AddElement(&m[i], &b[i]) + } +} + func (m MultiLinE6) Sum() extensions.E6 { s := m[0] for i := 1; i < len(m); i++ { diff --git a/field/koalabear/extensions/polynomial/multilin_e6_test.go b/field/koalabear/extensions/polynomial/multilin_e6_test.go index 8f3ffd46e5..e5a412f8a8 100644 --- a/field/koalabear/extensions/polynomial/multilin_e6_test.go +++ b/field/koalabear/extensions/polynomial/multilin_e6_test.go @@ -9,6 +9,7 @@ import ( "testing" "github.com/consensys/gnark-crypto/field/koalabear/extensions" + basepoly "github.com/consensys/gnark-crypto/field/koalabear/polynomial" "github.com/stretchr/testify/assert" ) @@ -83,3 +84,65 @@ func TestFoldedEqTableE6(t *testing.T) { } } + +func TestFoldFromBaseE6(t *testing.T) { + for _, n := range []int{2, 4, 8, 64} { + b := make(basepoly.MultiLin, n) + for i := range b { + b[i].MustSetRandom() + } + var r extensions.E6 + r.MustSetRandom() + + // reference: lift b to the extension and fold in place + lifted := make(MultiLinE6, n) + for i := range b { + lifted[i].SetElement(&b[i]) + } + lifted.Fold(r) + + var got MultiLinE6 + got.FoldFromBase(b, &r) + assert.Equal(t, lifted, got) + + // parallel version, split in two tasks, reusing a dirty destination + par := make(MultiLinE6, n) + task := par.FoldFromBaseParallel(b, &r) + mid := n / 2 + task(0, mid/2) + task(mid/2, mid) + assert.Equal(t, lifted, par) + } +} + +func TestEvaluateBaseE6(t *testing.T) { + for _, nbVars := range []int{0, 1, 2, 3, 6} { + b := make(basepoly.MultiLin, 1<= 0; i-- { + res.Mul(&res, v) + res.Add(&res, &(*p)[i]) + } + + return res +} + +// Clone returns a copy of the polynomial +func (p *Polynomial) Clone() Polynomial { + _p := make(Polynomial, len(*p)) + copy(_p, *p) + return _p +} + +// Set to another polynomial +func (p *Polynomial) Set(p1 Polynomial) { + if len(*p) != len(p1) { + *p = p1.Clone() + return + } + + for i := range len(p1) { + (*p)[i].Set(&p1[i]) + } +} + +// AddConstantInPlace adds a constant to the polynomial, modifying p +func (p *Polynomial) AddConstantInPlace(c *koalabear.Element) { + for i := range len(*p) { + (*p)[i].Add(&(*p)[i], c) + } +} + +// SubConstantInPlace subs a constant to the polynomial, modifying p +func (p *Polynomial) SubConstantInPlace(c *koalabear.Element) { + for i := range len(*p) { + (*p)[i].Sub(&(*p)[i], c) + } +} + +// ScaleInPlace multiplies p by v, modifying p +func (p *Polynomial) ScaleInPlace(c *koalabear.Element) { + for i := range len(*p) { + (*p)[i].Mul(&(*p)[i], c) + } +} + +// Scale multiplies p0 by v, storing the result in p +func (p *Polynomial) Scale(c *koalabear.Element, p0 Polynomial) { + if len(*p) != len(p0) { + *p = make(Polynomial, len(p0)) + } + for i := range len(p0) { + (*p)[i].Mul(c, &p0[i]) + } +} + +// Add adds p1 to p2 +// This function allocates a new slice unless p == p1 or p == p2 +func (p *Polynomial) Add(p1, p2 Polynomial) *Polynomial { + + bigger := p1 + smaller := p2 + if len(bigger) < len(smaller) { + bigger, smaller = smaller, bigger + } + + if len(*p) == len(bigger) && (&(*p)[0] == &bigger[0]) { + for i := range len(smaller) { + (*p)[i].Add(&(*p)[i], &smaller[i]) + } + return p + } + + if len(*p) == len(smaller) && (&(*p)[0] == &smaller[0]) { + for i := range len(smaller) { + (*p)[i].Add(&(*p)[i], &bigger[i]) + } + *p = append(*p, bigger[len(smaller):]...) + return p + } + + res := make(Polynomial, len(bigger)) + copy(res, bigger) + for i := range len(smaller) { + res[i].Add(&res[i], &smaller[i]) + } + *p = res + return p +} + +// Sub subtracts p2 from p1 +// TODO make interface more consistent with Add +func (p *Polynomial) Sub(p1, p2 Polynomial) *Polynomial { + if len(p1) != len(p2) || len(p2) != len(*p) { + return nil + } + for i := range len(*p) { + (*p)[i].Sub(&p1[i], &p2[i]) + } + return p +} + +// Equal checks equality between two polynomials +func (p *Polynomial) Equal(p1 Polynomial) bool { + if (*p == nil) != (p1 == nil) { + return false + } + + if len(*p) != len(p1) { + return false + } + + for i := range p1 { + if !(*p)[i].Equal(&p1[i]) { + return false + } + } + + return true +} + +func (p Polynomial) SetZero() { + for i := range len(p) { + p[i].SetZero() + } +} + +func (p Polynomial) Text(base int) string { + + var builder strings.Builder + + first := true + for d := len(p) - 1; d >= 0; d-- { + if p[d].IsZero() { + continue + } + + pD := p[d] + pDText := pD.Text(base) + + initialLen := builder.Len() + + if pDText[0] == '-' { + pDText = pDText[1:] + if first { + builder.WriteString("-") + } else { + builder.WriteString(" - ") + } + } else if !first { + builder.WriteString(" + ") + } + + first = false + + if !pD.IsOne() || d == 0 { + builder.WriteString(pDText) + } + + if builder.Len()-initialLen > 10 { + builder.WriteString("×") + } + + if d != 0 { + builder.WriteString("X") + } + if d > 1 { + builder.WriteString( + utils.ToSuperscript(strconv.Itoa(d)), + ) + } + + } + + if first { + return "0" + } + + return builder.String() +} + +// InterpolateOnRange maps vector v to polynomial f +// such that f(i) = v[i] for 0 ≤ i < len(v). +// len(f) = len(v) and deg(f) ≤ len(v) - 1 +func InterpolateOnRange(v []koalabear.Element) Polynomial { + nEvals := uint8(len(v)) + if int(nEvals) != len(v) { + panic("interpolation method too inefficient for nEvals > 255") + } + lagrange := getLagrangeBasis(nEvals) + + var res Polynomial + res.Scale(&v[0], lagrange[0]) + + temp := make(Polynomial, nEvals) + + for i := uint8(1); i < nEvals; i++ { + temp.Scale(&v[i], lagrange[i]) + res.Add(res, temp) + } + + return res +} + +// lagrange bases used by InterpolateOnRange +var lagrangeBasis sync.Map + +func getLagrangeBasis(domainSize uint8) []Polynomial { + if res, ok := lagrangeBasis.Load(domainSize); ok { + return res.([]Polynomial) + } + + // not found. compute + var res []Polynomial + if domainSize >= 2 { + res = computeLagrangeBasis(domainSize) + } else if domainSize == 1 { + res = []Polynomial{make(Polynomial, 1)} + res[0][0].SetOne() + } + lagrangeBasis.Store(domainSize, res) + + return res +} + +// computeLagrangeBasis precomputes in explicit coefficient form for each 0 ≤ l < domainSize the polynomial +// pₗ := X (X-1) ... (X-l-1) (X-l+1) ... (X - domainSize + 1) / ( l (l-1) ... 2 (-1) ... (l - domainSize +1) ) +// Note that pₗ(l) = 1 and pₗ(n) = 0 if 0 ≤ l < domainSize, n ≠ l +func computeLagrangeBasis(domainSize uint8) []Polynomial { + + constTerms := make([]koalabear.Element, domainSize) + for i := range domainSize { + constTerms[i].SetInt64(-int64(i)) + } + + res := make([]Polynomial, domainSize) + multScratch := make(Polynomial, domainSize-1) + + // compute pₗ + for l := range domainSize { + + // TODO @Tabaie Optimize this with some trees? O(log(domainSize)) polynomial mults instead of O(domainSize)? Then again it would be fewer big poly mults vs many small poly mults + d := uint8(0) //d is the current degree of res + for i := range domainSize { + if i == l { + continue + } + if d == 0 { + res[l] = make(Polynomial, domainSize) + res[l][domainSize-2] = constTerms[i] + res[l][domainSize-1].SetOne() + } else { + current := res[l][domainSize-d-2:] + timesConst := multScratch[domainSize-d-2:] + + timesConst.Scale(&constTerms[i], current[1:]) //TODO: Directly double and add since constTerms are tiny? (even less than 4 bits) + nonLeading := current[0 : d+1] + + nonLeading.Add(nonLeading, timesConst) + + } + d++ + } + + } + + // We have pₗ(i≠l)=0. Now scale so that pₗ(l)=1 + // Replace the constTerms with norms + for l := range domainSize { + constTerms[l].Neg(&constTerms[l]) + constTerms[l] = res[l].Eval(&constTerms[l]) + } + constTerms = koalabear.BatchInvert(constTerms) + for l := range domainSize { + res[l].ScaleInPlace(&constTerms[l]) + } + + return res +} diff --git a/field/koalabear/polynomial/polynomial_test.go b/field/koalabear/polynomial/polynomial_test.go new file mode 100644 index 0000000000..604af7aa72 --- /dev/null +++ b/field/koalabear/polynomial/polynomial_test.go @@ -0,0 +1,255 @@ +// Copyright 2020-2026 Consensys Software Inc. +// Licensed under the Apache License, Version 2.0. See the LICENSE file for details. + +// Code generated by consensys/gnark-crypto DO NOT EDIT + +package polynomial + +import ( + "math/big" + "testing" + + "github.com/consensys/gnark-crypto/field/koalabear" + "github.com/leanovate/gopter" + "github.com/leanovate/gopter/gen" + "github.com/leanovate/gopter/prop" + "github.com/stretchr/testify/assert" +) + +func TestPolynomialEval(t *testing.T) { + + // build polynomial + f := make(Polynomial, 20) + for i := range 20 { + f[i].SetOne() + } + + // random value + var point koalabear.Element + point.MustSetRandom() + + // compute manually f(val) + var expectedEval, one, den koalabear.Element + var expo big.Int + one.SetOne() + expo.SetUint64(20) + expectedEval.Exp(point, &expo). + Sub(&expectedEval, &one) + den.Sub(&point, &one) + expectedEval.Div(&expectedEval, &den) + + // compute purported evaluation + purportedEval := f.Eval(&point) + + // check + if !purportedEval.Equal(&expectedEval) { + t.Fatal("polynomial evaluation failed") + } +} + +func TestPolynomialAddConstantInPlace(t *testing.T) { + + // build polynomial + f := make(Polynomial, 20) + for i := range 20 { + f[i].SetOne() + } + + // constant to add + var c koalabear.Element + c.MustSetRandom() + + // add constant + f.AddConstantInPlace(&c) + + // check + var expectedCoeffs, one koalabear.Element + one.SetOne() + expectedCoeffs.Add(&one, &c) + for i := range 20 { + if !f[i].Equal(&expectedCoeffs) { + t.Fatal("AddConstantInPlace failed") + } + } +} + +func TestPolynomialSubConstantInPlace(t *testing.T) { + + // build polynomial + f := make(Polynomial, 20) + for i := range 20 { + f[i].SetOne() + } + + // constant to sub + var c koalabear.Element + c.MustSetRandom() + + // sub constant + f.SubConstantInPlace(&c) + + // check + var expectedCoeffs, one koalabear.Element + one.SetOne() + expectedCoeffs.Sub(&one, &c) + for i := range 20 { + if !f[i].Equal(&expectedCoeffs) { + t.Fatal("SubConstantInPlace failed") + } + } +} + +func TestPolynomialScaleInPlace(t *testing.T) { + + // build polynomial + f := make(Polynomial, 20) + for i := range 20 { + f[i].SetOne() + } + + // constant to scale by + var c koalabear.Element + c.MustSetRandom() + + // scale by constant + f.ScaleInPlace(&c) + + // check + for i := range 20 { + if !f[i].Equal(&c) { + t.Fatal("ScaleInPlace failed") + } + } + +} + +func TestPolynomialAdd(t *testing.T) { + + // build unbalanced polynomials + f1 := make(Polynomial, 20) + f1Backup := make(Polynomial, 20) + for i := range 20 { + f1[i].SetOne() + f1Backup[i].SetOne() + } + f2 := make(Polynomial, 10) + f2Backup := make(Polynomial, 10) + for i := range 10 { + f2[i].SetOne() + f2Backup[i].SetOne() + } + + // expected result + var one, two koalabear.Element + one.SetOne() + two.Double(&one) + expectedSum := make(Polynomial, 20) + for i := range 10 { + expectedSum[i].Set(&two) + } + for i := 10; i < 20; i++ { + expectedSum[i].Set(&one) + } + + // caller is empty + var g Polynomial + g.Add(f1, f2) + if !g.Equal(expectedSum) { + t.Fatal("add polynomials fails") + } + if !f1.Equal(f1Backup) { + t.Fatal("side effect, f1 should not have been modified") + } + if !f2.Equal(f2Backup) { + t.Fatal("side effect, f2 should not have been modified") + } + + // all operands are distinct + _f1 := f1.Clone() + _f1.Add(f1, f2) + if !_f1.Equal(expectedSum) { + t.Fatal("add polynomials fails") + } + if !f1.Equal(f1Backup) { + t.Fatal("side effect, f1 should not have been modified") + } + if !f2.Equal(f2Backup) { + t.Fatal("side effect, f2 should not have been modified") + } + + // first operand = caller + _f1 = f1.Clone() + _f2 := f2.Clone() + _f1.Add(_f1, _f2) + if !_f1.Equal(expectedSum) { + t.Fatal("add polynomials fails") + } + if !_f2.Equal(f2Backup) { + t.Fatal("side effect, _f2 should not have been modified") + } + + // second operand = caller + _f1 = f1.Clone() + _f2 = f2.Clone() + _f1.Add(_f2, _f1) + if !_f1.Equal(expectedSum) { + t.Fatal("add polynomials fails") + } + if !_f2.Equal(f2Backup) { + t.Fatal("side effect, _f2 should not have been modified") + } +} + +func TestPolynomialText(t *testing.T) { + var one, negTwo koalabear.Element + one.SetOne() + negTwo.SetInt64(-2) + + p := Polynomial{one, negTwo, one} + + assert.Equal(t, "X² - 2X + 1", p.Text(10)) +} + +func TestPrecomputeLagrange(t *testing.T) { + + testForDomainSize := func(domainSize uint8) bool { + polys := computeLagrangeBasis(domainSize) + + for l := range domainSize { + for i := range domainSize { + var I koalabear.Element + I.SetUint64(uint64(i)) + y := polys[l].Eval(&I) + + if i == l && !y.IsOne() || i != l && !y.IsZero() { + t.Errorf("domainSize = %d: p_%d(%d) = %s", domainSize, l, i, y.Text(10)) + return false + } + } + } + return true + } + + t.Parallel() + parameters := gopter.DefaultTestParameters() + + const maxLagrangeDomainSize = 12 + + parameters.MinSuccessfulTests = maxLagrangeDomainSize + + properties := gopter.NewProperties(parameters) + + properties.Property("l'th lagrange polynomials must evaluate to 1 on l and 0 on other values in the domain", prop.ForAll( + testForDomainSize, + gen.UInt8Range(2, maxLagrangeDomainSize), + )) + + properties.TestingRun(t, gopter.ConsoleReporter(false)) +} + +func TestLagrangeCache(t *testing.T) { + for _, i := range []int{5, 2, 8, 4, 6, 3, 0} { + b := getLagrangeBasis(uint8(i)) + assert.Equal(t, b, getLagrangeBasis(uint8(i))) // second call must yield the same result + } +} diff --git a/field/koalabear/polynomial/pool.go b/field/koalabear/polynomial/pool.go new file mode 100644 index 0000000000..be318bb1ad --- /dev/null +++ b/field/koalabear/polynomial/pool.go @@ -0,0 +1,191 @@ +// Copyright 2020-2026 Consensys Software Inc. +// Licensed under the Apache License, Version 2.0. See the LICENSE file for details. + +// Code generated by consensys/gnark-crypto DO NOT EDIT + +package polynomial + +import ( + "encoding/json" + "fmt" + "runtime" + "sort" + "sync" + "unsafe" + + "github.com/consensys/gnark-crypto/field/koalabear" +) + +// Memory management for polynomials +// WARNING: This is not thread safe TODO: Make sure that is not a problem +// TODO: There is a lot of "unsafe" memory management here and needs to be vetted thoroughly + +type sizedPool struct { + maxN int + pool sync.Pool + stats poolStats +} + +type inUseData struct { + allocatedFor []uintptr + pool *sizedPool +} + +type Pool struct { + //lock sync.Mutex + inUse sync.Map + subPools []sizedPool +} + +func (p *sizedPool) get(n int) *koalabear.Element { + p.stats.make(n) + return p.pool.Get().(*koalabear.Element) +} + +func (p *sizedPool) put(ptr *koalabear.Element) { + p.stats.dump() + p.pool.Put(ptr) +} + +func NewPool(maxN ...int) (pool Pool) { + + sort.Ints(maxN) + pool = Pool{ + subPools: make([]sizedPool, len(maxN)), + } + + for i := range pool.subPools { + subPool := &pool.subPools[i] + subPool.maxN = maxN[i] + subPool.pool = sync.Pool{ + New: func() any { + subPool.stats.Allocated++ + return getDataPointer(make([]koalabear.Element, 0, subPool.maxN)) + }, + } + } + return +} + +func (p *Pool) findCorrespondingPool(n int) *sizedPool { + poolI := 0 + for poolI < len(p.subPools) && n > p.subPools[poolI].maxN { + poolI++ + } + return &p.subPools[poolI] // out of bounds error here would mean that n is too large +} + +func (p *Pool) Make(n int) []koalabear.Element { + pool := p.findCorrespondingPool(n) + ptr := pool.get(n) + p.addInUse(ptr, pool) + return unsafe.Slice(ptr, n) +} + +// Dump dumps a set of polynomials into the pool +func (p *Pool) Dump(slices ...[]koalabear.Element) { + for _, slice := range slices { + ptr := getDataPointer(slice) + if metadata, ok := p.inUse.Load(ptr); ok { + p.inUse.Delete(ptr) + metadata.(inUseData).pool.put(ptr) + } else { + panic("attempting to dump a slice not created by the pool") + } + } +} + +func (p *Pool) addInUse(ptr *koalabear.Element, pool *sizedPool) { + pcs := make([]uintptr, 2) + n := runtime.Callers(3, pcs) + + if prevPcs, ok := p.inUse.Load(ptr); ok { // TODO: remove if unnecessary for security + panic(fmt.Errorf("re-allocated non-dumped slice, previously allocated at %v", runtime.CallersFrames(prevPcs.(inUseData).allocatedFor))) + } + p.inUse.Store(ptr, inUseData{ + allocatedFor: pcs[:n], + pool: pool, + }) +} + +func printFrame(frame runtime.Frame) { + fmt.Printf("\t%s line %d, function %s\n", frame.File, frame.Line, frame.Function) +} + +func (p *Pool) printInUse() { + fmt.Println("slices never dumped allocated at:") + p.inUse.Range(func(_, pcs any) bool { + fmt.Println("-------------------------") + + var frame runtime.Frame + frames := runtime.CallersFrames(pcs.(inUseData).allocatedFor) + more := true + for more { + frame, more = frames.Next() + printFrame(frame) + } + return true + }) +} + +type poolStats struct { + Used int + Allocated int + ReuseRate float64 + InUse int + GreatestNUsed int + SmallestNUsed int +} + +type poolsStats struct { + SubPools []poolStats + InUse int +} + +func (s *poolStats) make(n int) { + s.Used++ + s.InUse++ + if n > s.GreatestNUsed { + s.GreatestNUsed = n + } + if s.SmallestNUsed == 0 || s.SmallestNUsed > n { + s.SmallestNUsed = n + } +} + +func (s *poolStats) dump() { + s.InUse-- +} + +func (s *poolStats) finalize() { + s.ReuseRate = float64(s.Used) / float64(s.Allocated) +} + +func getDataPointer(slice []koalabear.Element) *koalabear.Element { + return (*koalabear.Element)(unsafe.SliceData(slice)) +} + +func (p *Pool) PrintPoolStats() { + InUse := 0 + subStats := make([]poolStats, len(p.subPools)) + for i := range p.subPools { + subPool := &p.subPools[i] + subPool.stats.finalize() + subStats[i] = subPool.stats + InUse += subPool.stats.InUse + } + + stats := poolsStats{ + SubPools: subStats, + InUse: InUse, + } + serialized, _ := json.MarshalIndent(stats, "", " ") + fmt.Println(string(serialized)) + p.printInUse() +} + +func (p *Pool) Clone(slice []koalabear.Element) []koalabear.Element { + res := p.Make(len(slice)) + copy(res, slice) + return res +} diff --git a/field/mamabear/extensions/e3.go b/field/mamabear/extensions/e3.go index baaac5d57e..748929b8f9 100644 --- a/field/mamabear/extensions/e3.go +++ b/field/mamabear/extensions/e3.go @@ -64,6 +64,32 @@ func (z *E3) Div(x, y *E3) *E3 { return z.Set(&r) } +// AddElement sets z to x + y, where y is an element of the base field embedded in E3, and returns z +func (z *E3) AddElement(x *E3, y *fr.Element) *E3 { + yc := *y + z.A0.Add(&x.A0, &yc) + z.A1 = x.A1 + z.A2 = x.A2 + return z +} + +// SubElement sets z to x - y, where y is an element of the base field embedded in E3, and returns z +func (z *E3) SubElement(x *E3, y *fr.Element) *E3 { + yc := *y + z.A0.Sub(&x.A0, &yc) + z.A1 = x.A1 + z.A2 = x.A2 + return z +} + +// SetElement sets z to x, an element of the base field embedded in E3, and returns z +func (z *E3) SetElement(x *fr.Element) *E3 { + v := *x + *z = E3{} + z.A0 = v + return z +} + // SetBigInt sets z to v (embedded via the unique ring homomorphism ℤ→𝔽, reduced modulo the base field order) and returns z func (z *E3) SetBigInt(v *big.Int) *E3 { *z = E3{} diff --git a/field/mamabear/extensions/e3_test.go b/field/mamabear/extensions/e3_test.go index c83c1f1f07..590b277321 100644 --- a/field/mamabear/extensions/e3_test.go +++ b/field/mamabear/extensions/e3_test.go @@ -543,3 +543,35 @@ func TestE3SetBytesCanonicalRejectsNonCanonical(t *testing.T) { require.True(t, y.Equal(&x), "E3 must be unchanged on error") } } + +func TestE3ElementOps(t *testing.T) { + for range 100 { + var x, lifted, got, want E3 + var e fr.Element + x.MustSetRandom() + e.MustSetRandom() + + var z E3 + z.SetElement(&e) + require.True(t, z.A0.Equal(&e)) + require.True(t, z.A1.IsZero()) + require.True(t, z.A2.IsZero()) + lifted = z + + got.AddElement(&x, &e) + want.Add(&x, &lifted) + require.True(t, got.Equal(&want), "AddElement") + + got.SubElement(&x, &e) + want.Sub(&x, &lifted) + require.True(t, got.Equal(&want), "SubElement") + + // aliasing of the receiver with the first operand + got.Set(&x) + got.AddElement(&got, &e) + want.Add(&x, &lifted) + require.True(t, got.Equal(&want), "AddElement alias") + got.SubElement(&got, &e) + require.True(t, got.Equal(&x), "SubElement alias") + } +} diff --git a/field/mamabear/polynomial/doc.go b/field/mamabear/polynomial/doc.go new file mode 100644 index 0000000000..aa346f3ea3 --- /dev/null +++ b/field/mamabear/polynomial/doc.go @@ -0,0 +1,7 @@ +// Copyright 2020-2026 Consensys Software Inc. +// Licensed under the Apache License, Version 2.0. See the LICENSE file for details. + +// Code generated by consensys/gnark-crypto DO NOT EDIT + +// Package polynomial provides polynomial methods and commitment schemes. +package polynomial diff --git a/field/mamabear/polynomial/multilin.go b/field/mamabear/polynomial/multilin.go new file mode 100644 index 0000000000..25171cd3d4 --- /dev/null +++ b/field/mamabear/polynomial/multilin.go @@ -0,0 +1,179 @@ +// Copyright 2020-2026 Consensys Software Inc. +// Licensed under the Apache License, Version 2.0. See the LICENSE file for details. + +// Code generated by consensys/gnark-crypto DO NOT EDIT + +package polynomial + +import ( + "math/bits" + + "github.com/consensys/gnark-crypto/field/mamabear" + "github.com/consensys/gnark-crypto/utils" +) + +// MultiLin tracks the values of a (dense i.e. not sparse) multilinear polynomial +// The variables are X₁ through Xₙ where n = log(len(.)) +// .[∑ᵢ 2ⁱ⁻¹ bₙ₋ᵢ] = the polynomial evaluated at (b₁, b₂, ..., bₙ) +// It is understood that any hypercube evaluation can be extrapolated to a multilinear polynomial +type MultiLin []mamabear.Element + +// Fold is partial evaluation function k[X₁, X₂, ..., Xₙ] → k[X₂, ..., Xₙ] by setting X₁=r +func (m *MultiLin) Fold(r mamabear.Element) { + mid := len(*m) / 2 + + bottom, top := (*m)[:mid], (*m)[mid:] + + var t mamabear.Element // no need to update the top part + + // updating bookkeeping table + // knowing that the polynomial f ∈ (k[X₂, ..., Xₙ])[X₁] is linear, we would get f(r) = f(0) + r(f(1) - f(0)) + // the following loop computes the evaluations of f(r) accordingly: + // f(r, b₂, ..., bₙ) = f(0, b₂, ..., bₙ) + r(f(1, b₂, ..., bₙ) - f(0, b₂, ..., bₙ)) + for i := range mid { + // table[i] ← table[i] + r (table[i + mid] - table[i]) + t.Sub(&top[i], &bottom[i]) + t.Mul(&t, &r) + bottom[i].Add(&bottom[i], &t) + } + + *m = (*m)[:mid] +} + +func (m *MultiLin) FoldParallel(r mamabear.Element) utils.Task { + mid := len(*m) / 2 + bottom, top := (*m)[:mid], (*m)[mid:] + + *m = bottom + + return func(start, end int) { + var t mamabear.Element // no need to update the top part + for i := start; i < end; i++ { + // table[i] ← table[i] + r (table[i + mid] - table[i]) + t.Sub(&top[i], &bottom[i]) + t.Mul(&t, &r) + bottom[i].Add(&bottom[i], &t) + } + } +} + +func (m MultiLin) Sum() mamabear.Element { + s := m[0] + for i := 1; i < len(m); i++ { + s.Add(&s, &m[i]) + } + return s +} + +func _clone(m MultiLin, p *Pool) MultiLin { + if p == nil { + return m.Clone() + } else { + return p.Clone(m) + } +} + +func _dump(m MultiLin, p *Pool) { + if p != nil { + p.Dump(m) + } +} + +// Evaluate extrapolate the value of the multilinear polynomial corresponding to m +// on the given coordinates +func (m MultiLin) Evaluate(coordinates []mamabear.Element, p *Pool) mamabear.Element { + // Folding is a mutating operation + bkCopy := _clone(m, p) + + // Evaluate step by step through repeated folding (i.e. evaluation at the first remaining variable) + for _, r := range coordinates { + bkCopy.Fold(r) + } + + result := bkCopy[0] + + _dump(bkCopy, p) + return result +} + +// Clone creates a deep copy of a bookkeeping table. +// Both multilinear interpolation and sumcheck require folding an underlying +// array, but folding changes the array. To do both one requires a deep copy +// of the bookkeeping table. +func (m MultiLin) Clone() MultiLin { + res := make(MultiLin, len(m)) + copy(res, m) + return res +} + +// Add two bookKeepingTables +func (m *MultiLin) Add(left, right MultiLin) { + size := len(left) + // Check that left and right have the same size + if len(right) != size || len(*m) != size { + panic("left, right and destination must have the right size") + } + + // Add elementwise + for i := range size { + (*m)[i].Add(&left[i], &right[i]) + } +} + +// EvalEq computes Eq(q₁, ... , qₙ, h₁, ... , hₙ) = Π₁ⁿ Eq(qᵢ, hᵢ) +// where Eq(x,y) = xy + (1-x)(1-y) = 1 - x - y + xy + xy interpolates +// +// _________________ +// | | | +// | 0 | 1 | +// |_______|_______| +// y | | | +// | 1 | 0 | +// |_______|_______| +// +// x +// +// In other words the polynomial evaluated here is the multilinear extrapolation of +// one that evaluates to q' == h' for vectors q', h' of binary values +func EvalEq(q, h []mamabear.Element) mamabear.Element { + var res, nxt, one, sum mamabear.Element + one.SetOne() + for i := range len(q) { + nxt.Mul(&q[i], &h[i]) // nxt <- qᵢ * hᵢ + nxt.Double(&nxt) // nxt <- 2 * qᵢ * hᵢ + nxt.Add(&nxt, &one) // nxt <- 1 + 2 * qᵢ * hᵢ + sum.Add(&q[i], &h[i]) // sum <- qᵢ + hᵢ TODO: Why not subtract one by one from nxt? More parallel? + + if i == 0 { + res.Sub(&nxt, &sum) // nxt <- 1 + 2 * qᵢ * hᵢ - qᵢ - hᵢ + } else { + nxt.Sub(&nxt, &sum) // nxt <- 1 + 2 * qᵢ * hᵢ - qᵢ - hᵢ + res.Mul(&res, &nxt) // res <- res * nxt + } + } + return res +} + +// Eq sets m to the representation of the polynomial Eq(q₁, ..., qₙ, *, ..., *) × m[0] +func (m *MultiLin) Eq(q []mamabear.Element) { + n := len(q) + + if len(*m) != 1<= 0; i-- { + res.Mul(&res, v) + res.Add(&res, &(*p)[i]) + } + + return res +} + +// Clone returns a copy of the polynomial +func (p *Polynomial) Clone() Polynomial { + _p := make(Polynomial, len(*p)) + copy(_p, *p) + return _p +} + +// Set to another polynomial +func (p *Polynomial) Set(p1 Polynomial) { + if len(*p) != len(p1) { + *p = p1.Clone() + return + } + + for i := range len(p1) { + (*p)[i].Set(&p1[i]) + } +} + +// AddConstantInPlace adds a constant to the polynomial, modifying p +func (p *Polynomial) AddConstantInPlace(c *mamabear.Element) { + for i := range len(*p) { + (*p)[i].Add(&(*p)[i], c) + } +} + +// SubConstantInPlace subs a constant to the polynomial, modifying p +func (p *Polynomial) SubConstantInPlace(c *mamabear.Element) { + for i := range len(*p) { + (*p)[i].Sub(&(*p)[i], c) + } +} + +// ScaleInPlace multiplies p by v, modifying p +func (p *Polynomial) ScaleInPlace(c *mamabear.Element) { + for i := range len(*p) { + (*p)[i].Mul(&(*p)[i], c) + } +} + +// Scale multiplies p0 by v, storing the result in p +func (p *Polynomial) Scale(c *mamabear.Element, p0 Polynomial) { + if len(*p) != len(p0) { + *p = make(Polynomial, len(p0)) + } + for i := range len(p0) { + (*p)[i].Mul(c, &p0[i]) + } +} + +// Add adds p1 to p2 +// This function allocates a new slice unless p == p1 or p == p2 +func (p *Polynomial) Add(p1, p2 Polynomial) *Polynomial { + + bigger := p1 + smaller := p2 + if len(bigger) < len(smaller) { + bigger, smaller = smaller, bigger + } + + if len(*p) == len(bigger) && (&(*p)[0] == &bigger[0]) { + for i := range len(smaller) { + (*p)[i].Add(&(*p)[i], &smaller[i]) + } + return p + } + + if len(*p) == len(smaller) && (&(*p)[0] == &smaller[0]) { + for i := range len(smaller) { + (*p)[i].Add(&(*p)[i], &bigger[i]) + } + *p = append(*p, bigger[len(smaller):]...) + return p + } + + res := make(Polynomial, len(bigger)) + copy(res, bigger) + for i := range len(smaller) { + res[i].Add(&res[i], &smaller[i]) + } + *p = res + return p +} + +// Sub subtracts p2 from p1 +// TODO make interface more consistent with Add +func (p *Polynomial) Sub(p1, p2 Polynomial) *Polynomial { + if len(p1) != len(p2) || len(p2) != len(*p) { + return nil + } + for i := range len(*p) { + (*p)[i].Sub(&p1[i], &p2[i]) + } + return p +} + +// Equal checks equality between two polynomials +func (p *Polynomial) Equal(p1 Polynomial) bool { + if (*p == nil) != (p1 == nil) { + return false + } + + if len(*p) != len(p1) { + return false + } + + for i := range p1 { + if !(*p)[i].Equal(&p1[i]) { + return false + } + } + + return true +} + +func (p Polynomial) SetZero() { + for i := range len(p) { + p[i].SetZero() + } +} + +func (p Polynomial) Text(base int) string { + + var builder strings.Builder + + first := true + for d := len(p) - 1; d >= 0; d-- { + if p[d].IsZero() { + continue + } + + pD := p[d] + pDText := pD.Text(base) + + initialLen := builder.Len() + + if pDText[0] == '-' { + pDText = pDText[1:] + if first { + builder.WriteString("-") + } else { + builder.WriteString(" - ") + } + } else if !first { + builder.WriteString(" + ") + } + + first = false + + if !pD.IsOne() || d == 0 { + builder.WriteString(pDText) + } + + if builder.Len()-initialLen > 10 { + builder.WriteString("×") + } + + if d != 0 { + builder.WriteString("X") + } + if d > 1 { + builder.WriteString( + utils.ToSuperscript(strconv.Itoa(d)), + ) + } + + } + + if first { + return "0" + } + + return builder.String() +} + +// InterpolateOnRange maps vector v to polynomial f +// such that f(i) = v[i] for 0 ≤ i < len(v). +// len(f) = len(v) and deg(f) ≤ len(v) - 1 +func InterpolateOnRange(v []mamabear.Element) Polynomial { + nEvals := uint8(len(v)) + if int(nEvals) != len(v) { + panic("interpolation method too inefficient for nEvals > 255") + } + lagrange := getLagrangeBasis(nEvals) + + var res Polynomial + res.Scale(&v[0], lagrange[0]) + + temp := make(Polynomial, nEvals) + + for i := uint8(1); i < nEvals; i++ { + temp.Scale(&v[i], lagrange[i]) + res.Add(res, temp) + } + + return res +} + +// lagrange bases used by InterpolateOnRange +var lagrangeBasis sync.Map + +func getLagrangeBasis(domainSize uint8) []Polynomial { + if res, ok := lagrangeBasis.Load(domainSize); ok { + return res.([]Polynomial) + } + + // not found. compute + var res []Polynomial + if domainSize >= 2 { + res = computeLagrangeBasis(domainSize) + } else if domainSize == 1 { + res = []Polynomial{make(Polynomial, 1)} + res[0][0].SetOne() + } + lagrangeBasis.Store(domainSize, res) + + return res +} + +// computeLagrangeBasis precomputes in explicit coefficient form for each 0 ≤ l < domainSize the polynomial +// pₗ := X (X-1) ... (X-l-1) (X-l+1) ... (X - domainSize + 1) / ( l (l-1) ... 2 (-1) ... (l - domainSize +1) ) +// Note that pₗ(l) = 1 and pₗ(n) = 0 if 0 ≤ l < domainSize, n ≠ l +func computeLagrangeBasis(domainSize uint8) []Polynomial { + + constTerms := make([]mamabear.Element, domainSize) + for i := range domainSize { + constTerms[i].SetInt64(-int64(i)) + } + + res := make([]Polynomial, domainSize) + multScratch := make(Polynomial, domainSize-1) + + // compute pₗ + for l := range domainSize { + + // TODO @Tabaie Optimize this with some trees? O(log(domainSize)) polynomial mults instead of O(domainSize)? Then again it would be fewer big poly mults vs many small poly mults + d := uint8(0) //d is the current degree of res + for i := range domainSize { + if i == l { + continue + } + if d == 0 { + res[l] = make(Polynomial, domainSize) + res[l][domainSize-2] = constTerms[i] + res[l][domainSize-1].SetOne() + } else { + current := res[l][domainSize-d-2:] + timesConst := multScratch[domainSize-d-2:] + + timesConst.Scale(&constTerms[i], current[1:]) //TODO: Directly double and add since constTerms are tiny? (even less than 4 bits) + nonLeading := current[0 : d+1] + + nonLeading.Add(nonLeading, timesConst) + + } + d++ + } + + } + + // We have pₗ(i≠l)=0. Now scale so that pₗ(l)=1 + // Replace the constTerms with norms + for l := range domainSize { + constTerms[l].Neg(&constTerms[l]) + constTerms[l] = res[l].Eval(&constTerms[l]) + } + constTerms = mamabear.BatchInvert(constTerms) + for l := range domainSize { + res[l].ScaleInPlace(&constTerms[l]) + } + + return res +} diff --git a/field/mamabear/polynomial/polynomial_test.go b/field/mamabear/polynomial/polynomial_test.go new file mode 100644 index 0000000000..745b630258 --- /dev/null +++ b/field/mamabear/polynomial/polynomial_test.go @@ -0,0 +1,255 @@ +// Copyright 2020-2026 Consensys Software Inc. +// Licensed under the Apache License, Version 2.0. See the LICENSE file for details. + +// Code generated by consensys/gnark-crypto DO NOT EDIT + +package polynomial + +import ( + "math/big" + "testing" + + "github.com/consensys/gnark-crypto/field/mamabear" + "github.com/leanovate/gopter" + "github.com/leanovate/gopter/gen" + "github.com/leanovate/gopter/prop" + "github.com/stretchr/testify/assert" +) + +func TestPolynomialEval(t *testing.T) { + + // build polynomial + f := make(Polynomial, 20) + for i := range 20 { + f[i].SetOne() + } + + // random value + var point mamabear.Element + point.MustSetRandom() + + // compute manually f(val) + var expectedEval, one, den mamabear.Element + var expo big.Int + one.SetOne() + expo.SetUint64(20) + expectedEval.Exp(point, &expo). + Sub(&expectedEval, &one) + den.Sub(&point, &one) + expectedEval.Div(&expectedEval, &den) + + // compute purported evaluation + purportedEval := f.Eval(&point) + + // check + if !purportedEval.Equal(&expectedEval) { + t.Fatal("polynomial evaluation failed") + } +} + +func TestPolynomialAddConstantInPlace(t *testing.T) { + + // build polynomial + f := make(Polynomial, 20) + for i := range 20 { + f[i].SetOne() + } + + // constant to add + var c mamabear.Element + c.MustSetRandom() + + // add constant + f.AddConstantInPlace(&c) + + // check + var expectedCoeffs, one mamabear.Element + one.SetOne() + expectedCoeffs.Add(&one, &c) + for i := range 20 { + if !f[i].Equal(&expectedCoeffs) { + t.Fatal("AddConstantInPlace failed") + } + } +} + +func TestPolynomialSubConstantInPlace(t *testing.T) { + + // build polynomial + f := make(Polynomial, 20) + for i := range 20 { + f[i].SetOne() + } + + // constant to sub + var c mamabear.Element + c.MustSetRandom() + + // sub constant + f.SubConstantInPlace(&c) + + // check + var expectedCoeffs, one mamabear.Element + one.SetOne() + expectedCoeffs.Sub(&one, &c) + for i := range 20 { + if !f[i].Equal(&expectedCoeffs) { + t.Fatal("SubConstantInPlace failed") + } + } +} + +func TestPolynomialScaleInPlace(t *testing.T) { + + // build polynomial + f := make(Polynomial, 20) + for i := range 20 { + f[i].SetOne() + } + + // constant to scale by + var c mamabear.Element + c.MustSetRandom() + + // scale by constant + f.ScaleInPlace(&c) + + // check + for i := range 20 { + if !f[i].Equal(&c) { + t.Fatal("ScaleInPlace failed") + } + } + +} + +func TestPolynomialAdd(t *testing.T) { + + // build unbalanced polynomials + f1 := make(Polynomial, 20) + f1Backup := make(Polynomial, 20) + for i := range 20 { + f1[i].SetOne() + f1Backup[i].SetOne() + } + f2 := make(Polynomial, 10) + f2Backup := make(Polynomial, 10) + for i := range 10 { + f2[i].SetOne() + f2Backup[i].SetOne() + } + + // expected result + var one, two mamabear.Element + one.SetOne() + two.Double(&one) + expectedSum := make(Polynomial, 20) + for i := range 10 { + expectedSum[i].Set(&two) + } + for i := 10; i < 20; i++ { + expectedSum[i].Set(&one) + } + + // caller is empty + var g Polynomial + g.Add(f1, f2) + if !g.Equal(expectedSum) { + t.Fatal("add polynomials fails") + } + if !f1.Equal(f1Backup) { + t.Fatal("side effect, f1 should not have been modified") + } + if !f2.Equal(f2Backup) { + t.Fatal("side effect, f2 should not have been modified") + } + + // all operands are distinct + _f1 := f1.Clone() + _f1.Add(f1, f2) + if !_f1.Equal(expectedSum) { + t.Fatal("add polynomials fails") + } + if !f1.Equal(f1Backup) { + t.Fatal("side effect, f1 should not have been modified") + } + if !f2.Equal(f2Backup) { + t.Fatal("side effect, f2 should not have been modified") + } + + // first operand = caller + _f1 = f1.Clone() + _f2 := f2.Clone() + _f1.Add(_f1, _f2) + if !_f1.Equal(expectedSum) { + t.Fatal("add polynomials fails") + } + if !_f2.Equal(f2Backup) { + t.Fatal("side effect, _f2 should not have been modified") + } + + // second operand = caller + _f1 = f1.Clone() + _f2 = f2.Clone() + _f1.Add(_f2, _f1) + if !_f1.Equal(expectedSum) { + t.Fatal("add polynomials fails") + } + if !_f2.Equal(f2Backup) { + t.Fatal("side effect, _f2 should not have been modified") + } +} + +func TestPolynomialText(t *testing.T) { + var one, negTwo mamabear.Element + one.SetOne() + negTwo.SetInt64(-2) + + p := Polynomial{one, negTwo, one} + + assert.Equal(t, "X² - 2X + 1", p.Text(10)) +} + +func TestPrecomputeLagrange(t *testing.T) { + + testForDomainSize := func(domainSize uint8) bool { + polys := computeLagrangeBasis(domainSize) + + for l := range domainSize { + for i := range domainSize { + var I mamabear.Element + I.SetUint64(uint64(i)) + y := polys[l].Eval(&I) + + if i == l && !y.IsOne() || i != l && !y.IsZero() { + t.Errorf("domainSize = %d: p_%d(%d) = %s", domainSize, l, i, y.Text(10)) + return false + } + } + } + return true + } + + t.Parallel() + parameters := gopter.DefaultTestParameters() + + const maxLagrangeDomainSize = 12 + + parameters.MinSuccessfulTests = maxLagrangeDomainSize + + properties := gopter.NewProperties(parameters) + + properties.Property("l'th lagrange polynomials must evaluate to 1 on l and 0 on other values in the domain", prop.ForAll( + testForDomainSize, + gen.UInt8Range(2, maxLagrangeDomainSize), + )) + + properties.TestingRun(t, gopter.ConsoleReporter(false)) +} + +func TestLagrangeCache(t *testing.T) { + for _, i := range []int{5, 2, 8, 4, 6, 3, 0} { + b := getLagrangeBasis(uint8(i)) + assert.Equal(t, b, getLagrangeBasis(uint8(i))) // second call must yield the same result + } +} diff --git a/field/mamabear/polynomial/pool.go b/field/mamabear/polynomial/pool.go new file mode 100644 index 0000000000..48e5ec1bf2 --- /dev/null +++ b/field/mamabear/polynomial/pool.go @@ -0,0 +1,191 @@ +// Copyright 2020-2026 Consensys Software Inc. +// Licensed under the Apache License, Version 2.0. See the LICENSE file for details. + +// Code generated by consensys/gnark-crypto DO NOT EDIT + +package polynomial + +import ( + "encoding/json" + "fmt" + "runtime" + "sort" + "sync" + "unsafe" + + "github.com/consensys/gnark-crypto/field/mamabear" +) + +// Memory management for polynomials +// WARNING: This is not thread safe TODO: Make sure that is not a problem +// TODO: There is a lot of "unsafe" memory management here and needs to be vetted thoroughly + +type sizedPool struct { + maxN int + pool sync.Pool + stats poolStats +} + +type inUseData struct { + allocatedFor []uintptr + pool *sizedPool +} + +type Pool struct { + //lock sync.Mutex + inUse sync.Map + subPools []sizedPool +} + +func (p *sizedPool) get(n int) *mamabear.Element { + p.stats.make(n) + return p.pool.Get().(*mamabear.Element) +} + +func (p *sizedPool) put(ptr *mamabear.Element) { + p.stats.dump() + p.pool.Put(ptr) +} + +func NewPool(maxN ...int) (pool Pool) { + + sort.Ints(maxN) + pool = Pool{ + subPools: make([]sizedPool, len(maxN)), + } + + for i := range pool.subPools { + subPool := &pool.subPools[i] + subPool.maxN = maxN[i] + subPool.pool = sync.Pool{ + New: func() any { + subPool.stats.Allocated++ + return getDataPointer(make([]mamabear.Element, 0, subPool.maxN)) + }, + } + } + return +} + +func (p *Pool) findCorrespondingPool(n int) *sizedPool { + poolI := 0 + for poolI < len(p.subPools) && n > p.subPools[poolI].maxN { + poolI++ + } + return &p.subPools[poolI] // out of bounds error here would mean that n is too large +} + +func (p *Pool) Make(n int) []mamabear.Element { + pool := p.findCorrespondingPool(n) + ptr := pool.get(n) + p.addInUse(ptr, pool) + return unsafe.Slice(ptr, n) +} + +// Dump dumps a set of polynomials into the pool +func (p *Pool) Dump(slices ...[]mamabear.Element) { + for _, slice := range slices { + ptr := getDataPointer(slice) + if metadata, ok := p.inUse.Load(ptr); ok { + p.inUse.Delete(ptr) + metadata.(inUseData).pool.put(ptr) + } else { + panic("attempting to dump a slice not created by the pool") + } + } +} + +func (p *Pool) addInUse(ptr *mamabear.Element, pool *sizedPool) { + pcs := make([]uintptr, 2) + n := runtime.Callers(3, pcs) + + if prevPcs, ok := p.inUse.Load(ptr); ok { // TODO: remove if unnecessary for security + panic(fmt.Errorf("re-allocated non-dumped slice, previously allocated at %v", runtime.CallersFrames(prevPcs.(inUseData).allocatedFor))) + } + p.inUse.Store(ptr, inUseData{ + allocatedFor: pcs[:n], + pool: pool, + }) +} + +func printFrame(frame runtime.Frame) { + fmt.Printf("\t%s line %d, function %s\n", frame.File, frame.Line, frame.Function) +} + +func (p *Pool) printInUse() { + fmt.Println("slices never dumped allocated at:") + p.inUse.Range(func(_, pcs any) bool { + fmt.Println("-------------------------") + + var frame runtime.Frame + frames := runtime.CallersFrames(pcs.(inUseData).allocatedFor) + more := true + for more { + frame, more = frames.Next() + printFrame(frame) + } + return true + }) +} + +type poolStats struct { + Used int + Allocated int + ReuseRate float64 + InUse int + GreatestNUsed int + SmallestNUsed int +} + +type poolsStats struct { + SubPools []poolStats + InUse int +} + +func (s *poolStats) make(n int) { + s.Used++ + s.InUse++ + if n > s.GreatestNUsed { + s.GreatestNUsed = n + } + if s.SmallestNUsed == 0 || s.SmallestNUsed > n { + s.SmallestNUsed = n + } +} + +func (s *poolStats) dump() { + s.InUse-- +} + +func (s *poolStats) finalize() { + s.ReuseRate = float64(s.Used) / float64(s.Allocated) +} + +func getDataPointer(slice []mamabear.Element) *mamabear.Element { + return (*mamabear.Element)(unsafe.SliceData(slice)) +} + +func (p *Pool) PrintPoolStats() { + InUse := 0 + subStats := make([]poolStats, len(p.subPools)) + for i := range p.subPools { + subPool := &p.subPools[i] + subPool.stats.finalize() + subStats[i] = subPool.stats + InUse += subPool.stats.InUse + } + + stats := poolsStats{ + SubPools: subStats, + InUse: InUse, + } + serialized, _ := json.MarshalIndent(stats, "", " ") + fmt.Println(string(serialized)) + p.printInUse() +} + +func (p *Pool) Clone(slice []mamabear.Element) []mamabear.Element { + res := p.Make(len(slice)) + copy(res, slice) + return res +} diff --git a/internal/generator/field/config/field_config.go b/internal/generator/field/config/field_config.go index 4021921d4d..60092eada5 100644 --- a/internal/generator/field/config/field_config.go +++ b/internal/generator/field/config/field_config.go @@ -939,6 +939,8 @@ type FieldDependency struct { FieldPackagePath string FieldPackageName string ExtensionDegree int + // BaseFieldPackagePath is the import path of the base field of an extension + BaseFieldPackagePath string } // ExtensionName returns the name of the extension (e.g. "E6"), or the empty diff --git a/internal/generator/field/template/extensions/e2.go.tmpl b/internal/generator/field/template/extensions/e2.go.tmpl index 4659d35f11..3032902239 100644 --- a/internal/generator/field/template/extensions/e2.go.tmpl +++ b/internal/generator/field/template/extensions/e2.go.tmpl @@ -81,6 +81,30 @@ func (z *E2) SetUint64(v uint64) *E2 { return z } +// AddElement sets z to x + y, where y is an element of the base field embedded in E2, and returns z +func (z *E2) AddElement(x *E2, y *fr.Element) *E2 { + yc := *y + z.A0.Add(&x.A0, &yc) + z.A1 = x.A1 + return z +} + +// SubElement sets z to x - y, where y is an element of the base field embedded in E2, and returns z +func (z *E2) SubElement(x *E2, y *fr.Element) *E2 { + yc := *y + z.A0.Sub(&x.A0, &yc) + z.A1 = x.A1 + return z +} + +// SetElement sets z to x, an element of the base field embedded in E2, and returns z +func (z *E2) SetElement(x *fr.Element) *E2 { + v := *x + *z = E2{} + z.A0 = v + return z +} + // SetBigInt sets z to v (embedded via the unique ring homomorphism ℤ→𝔽, reduced modulo the base field order) and returns z func (z *E2) SetBigInt(v *big.Int) *E2 { *z = E2{} diff --git a/internal/generator/field/template/extensions/e2_test.go.tmpl b/internal/generator/field/template/extensions/e2_test.go.tmpl index 84398b3bec..ef6a39367d 100644 --- a/internal/generator/field/template/extensions/e2_test.go.tmpl +++ b/internal/generator/field/template/extensions/e2_test.go.tmpl @@ -636,3 +636,34 @@ func TestE2SetBytesCanonicalRejectsNonCanonical(t *testing.T) { require.True(t, y.Equal(&x), "E2 must be unchanged on error") } } + +func TestE2ElementOps(t *testing.T) { + for range 100 { + var x, lifted, got, want E2 + var e fr.Element + x.MustSetRandom() + e.MustSetRandom() + + var z E2 + z.SetElement(&e) + require.True(t, z.A0.Equal(&e)) + require.True(t, z.A1.IsZero()) + lifted = z + + got.AddElement(&x, &e) + want.Add(&x, &lifted) + require.True(t, got.Equal(&want), "AddElement") + + got.SubElement(&x, &e) + want.Sub(&x, &lifted) + require.True(t, got.Equal(&want), "SubElement") + + // aliasing of the receiver with the first operand + got.Set(&x) + got.AddElement(&got, &e) + want.Add(&x, &lifted) + require.True(t, got.Equal(&want), "AddElement alias") + got.SubElement(&got, &e) + require.True(t, got.Equal(&x), "SubElement alias") + } +} diff --git a/internal/generator/field/template/extensions/e3.go.tmpl b/internal/generator/field/template/extensions/e3.go.tmpl index 959ee05184..b5f9e31ed8 100644 --- a/internal/generator/field/template/extensions/e3.go.tmpl +++ b/internal/generator/field/template/extensions/e3.go.tmpl @@ -56,6 +56,32 @@ func (z *E3) Div(x, y *E3) *E3 { return z.Set(&r) } +// AddElement sets z to x + y, where y is an element of the base field embedded in E3, and returns z +func (z *E3) AddElement(x *E3, y *fr.Element) *E3 { + yc := *y + z.A0.Add(&x.A0, &yc) + z.A1 = x.A1 + z.A2 = x.A2 + return z +} + +// SubElement sets z to x - y, where y is an element of the base field embedded in E3, and returns z +func (z *E3) SubElement(x *E3, y *fr.Element) *E3 { + yc := *y + z.A0.Sub(&x.A0, &yc) + z.A1 = x.A1 + z.A2 = x.A2 + return z +} + +// SetElement sets z to x, an element of the base field embedded in E3, and returns z +func (z *E3) SetElement(x *fr.Element) *E3 { + v := *x + *z = E3{} + z.A0 = v + return z +} + // SetBigInt sets z to v (embedded via the unique ring homomorphism ℤ→𝔽, reduced modulo the base field order) and returns z func (z *E3) SetBigInt(v *big.Int) *E3 { *z = E3{} diff --git a/internal/generator/field/template/extensions/e3_test.go.tmpl b/internal/generator/field/template/extensions/e3_test.go.tmpl index 83f5fe8cda..72c9878191 100644 --- a/internal/generator/field/template/extensions/e3_test.go.tmpl +++ b/internal/generator/field/template/extensions/e3_test.go.tmpl @@ -534,3 +534,35 @@ func TestE3SetBytesCanonicalRejectsNonCanonical(t *testing.T) { require.True(t, y.Equal(&x), "E3 must be unchanged on error") } } + +func TestE3ElementOps(t *testing.T) { + for range 100 { + var x, lifted, got, want E3 + var e fr.Element + x.MustSetRandom() + e.MustSetRandom() + + var z E3 + z.SetElement(&e) + require.True(t, z.A0.Equal(&e)) + require.True(t, z.A1.IsZero()) + require.True(t, z.A2.IsZero()) + lifted = z + + got.AddElement(&x, &e) + want.Add(&x, &lifted) + require.True(t, got.Equal(&want), "AddElement") + + got.SubElement(&x, &e) + want.Sub(&x, &lifted) + require.True(t, got.Equal(&want), "SubElement") + + // aliasing of the receiver with the first operand + got.Set(&x) + got.AddElement(&got, &e) + want.Add(&x, &lifted) + require.True(t, got.Equal(&want), "AddElement alias") + got.SubElement(&got, &e) + require.True(t, got.Equal(&x), "SubElement alias") + } +} diff --git a/internal/generator/field/template/extensions/e4.go.tmpl b/internal/generator/field/template/extensions/e4.go.tmpl index 766e58d57c..dd7fd2633a 100644 --- a/internal/generator/field/template/extensions/e4.go.tmpl +++ b/internal/generator/field/template/extensions/e4.go.tmpl @@ -92,6 +92,34 @@ func (z *E4) SetUint64(v uint64) *E4 { return z } +// AddElement sets z to x + y, where y is an element of the base field embedded in E4, and returns z +func (z *E4) AddElement(x *E4, y *fr.Element) *E4 { + yc := *y + z.B0.A0.Add(&x.B0.A0, &yc) + z.B0.A1 = x.B0.A1 + z.B1.A0 = x.B1.A0 + z.B1.A1 = x.B1.A1 + return z +} + +// SubElement sets z to x - y, where y is an element of the base field embedded in E4, and returns z +func (z *E4) SubElement(x *E4, y *fr.Element) *E4 { + yc := *y + z.B0.A0.Sub(&x.B0.A0, &yc) + z.B0.A1 = x.B0.A1 + z.B1.A0 = x.B1.A0 + z.B1.A1 = x.B1.A1 + return z +} + +// SetElement sets z to x, an element of the base field embedded in E4, and returns z +func (z *E4) SetElement(x *fr.Element) *E4 { + v := *x + *z = E4{} + z.B0.A0 = v + return z +} + // SetBigInt sets z to v (embedded via the unique ring homomorphism ℤ→𝔽, reduced modulo the base field order) and returns z func (z *E4) SetBigInt(v *big.Int) *E4 { *z = E4{} diff --git a/internal/generator/field/template/extensions/e4_test.go.tmpl b/internal/generator/field/template/extensions/e4_test.go.tmpl index dcf32e9daf..174c0702e9 100644 --- a/internal/generator/field/template/extensions/e4_test.go.tmpl +++ b/internal/generator/field/template/extensions/e4_test.go.tmpl @@ -1204,3 +1204,36 @@ func TestE4SetBytesCanonicalRejectsNonCanonical(t *testing.T) { require.True(t, y.Equal(&x), "E4 must be unchanged on error") } } + +func TestE4ElementOps(t *testing.T) { + for range 100 { + var x, lifted, got, want E4 + var e fr.Element + x.MustSetRandom() + e.MustSetRandom() + + var z E4 + z.SetElement(&e) + require.True(t, z.B0.A0.Equal(&e)) + require.True(t, z.B0.A1.IsZero()) + require.True(t, z.B1.A0.IsZero()) + require.True(t, z.B1.A1.IsZero()) + lifted = z + + got.AddElement(&x, &e) + want.Add(&x, &lifted) + require.True(t, got.Equal(&want), "AddElement") + + got.SubElement(&x, &e) + want.Sub(&x, &lifted) + require.True(t, got.Equal(&want), "SubElement") + + // aliasing of the receiver with the first operand + got.Set(&x) + got.AddElement(&got, &e) + want.Add(&x, &lifted) + require.True(t, got.Equal(&want), "AddElement alias") + got.SubElement(&got, &e) + require.True(t, got.Equal(&x), "SubElement alias") + } +} diff --git a/internal/generator/field/template/extensions/e6.go.tmpl b/internal/generator/field/template/extensions/e6.go.tmpl index 7f996bbcc2..8a87423a23 100644 --- a/internal/generator/field/template/extensions/e6.go.tmpl +++ b/internal/generator/field/template/extensions/e6.go.tmpl @@ -85,6 +85,38 @@ func (z *E6) SetUint64(v uint64) *E6 { return z } +// AddElement sets z to x + y, where y is an element of the base field embedded in E6, and returns z +func (z *E6) AddElement(x *E6, y *fr.Element) *E6 { + yc := *y + z.B0.A0.Add(&x.B0.A0, &yc) + z.B0.A1 = x.B0.A1 + z.B1.A0 = x.B1.A0 + z.B1.A1 = x.B1.A1 + z.B2.A0 = x.B2.A0 + z.B2.A1 = x.B2.A1 + return z +} + +// SubElement sets z to x - y, where y is an element of the base field embedded in E6, and returns z +func (z *E6) SubElement(x *E6, y *fr.Element) *E6 { + yc := *y + z.B0.A0.Sub(&x.B0.A0, &yc) + z.B0.A1 = x.B0.A1 + z.B1.A0 = x.B1.A0 + z.B1.A1 = x.B1.A1 + z.B2.A0 = x.B2.A0 + z.B2.A1 = x.B2.A1 + return z +} + +// SetElement sets z to x, an element of the base field embedded in E6, and returns z +func (z *E6) SetElement(x *fr.Element) *E6 { + v := *x + *z = E6{} + z.B0.A0 = v + return z +} + // SetBigInt sets z to v (embedded via the unique ring homomorphism ℤ→𝔽, reduced modulo the base field order) and returns z func (z *E6) SetBigInt(v *big.Int) *E6 { *z = E6{} diff --git a/internal/generator/field/template/extensions/e6_test.go.tmpl b/internal/generator/field/template/extensions/e6_test.go.tmpl index d6c4fdcadc..e7f0d8b35b 100644 --- a/internal/generator/field/template/extensions/e6_test.go.tmpl +++ b/internal/generator/field/template/extensions/e6_test.go.tmpl @@ -399,3 +399,38 @@ func TestE6SetBytesCanonicalRejectsNonCanonical(t *testing.T) { require.True(t, y.Equal(&x), "E6 must be unchanged on error") } } + +func TestE6ElementOps(t *testing.T) { + for range 100 { + var x, lifted, got, want E6 + var e fr.Element + x.MustSetRandom() + e.MustSetRandom() + + var z E6 + z.SetElement(&e) + require.True(t, z.B0.A0.Equal(&e)) + require.True(t, z.B0.A1.IsZero()) + require.True(t, z.B1.A0.IsZero()) + require.True(t, z.B1.A1.IsZero()) + require.True(t, z.B2.A0.IsZero()) + require.True(t, z.B2.A1.IsZero()) + lifted = z + + got.AddElement(&x, &e) + want.Add(&x, &lifted) + require.True(t, got.Equal(&want), "AddElement") + + got.SubElement(&x, &e) + want.Sub(&x, &lifted) + require.True(t, got.Equal(&want), "SubElement") + + // aliasing of the receiver with the first operand + got.Set(&x) + got.AddElement(&got, &e) + want.Add(&x, &lifted) + require.True(t, got.Equal(&want), "AddElement alias") + got.SubElement(&got, &e) + require.True(t, got.Equal(&x), "SubElement alias") + } +} diff --git a/internal/generator/main.go b/internal/generator/main.go index eb30909e29..4e4129b879 100644 --- a/internal/generator/main.go +++ b/internal/generator/main.go @@ -91,14 +91,22 @@ func main() { field.WithIOP(), )) + // polynomial package (Polynomial, MultiLin, Pool, ...) over the base field + assertNoError(polynomial.Generate(fieldConfig.FieldDependency{ + FieldPackagePath: "github.com/consensys/gnark-crypto/field/" + f.Name, + FieldPackageName: f.Name, + ElementType: f.Name + ".Element", + }, filepath.Join(outputDir, "polynomial"), true, true, gen)) + // polynomial package (Polynomial, MultiLin, Pool, ...) over the // requested extensions for i, degree := range f.PolynomialExtensions { extInfo := fieldConfig.FieldDependency{ - FieldPackagePath: "github.com/consensys/gnark-crypto/field/" + f.Name + "/extensions", - FieldPackageName: "extensions", - ElementType: fmt.Sprintf("extensions.E%d", degree), - ExtensionDegree: degree, + FieldPackagePath: "github.com/consensys/gnark-crypto/field/" + f.Name + "/extensions", + FieldPackageName: "extensions", + ElementType: fmt.Sprintf("extensions.E%d", degree), + ExtensionDegree: degree, + BaseFieldPackagePath: "github.com/consensys/gnark-crypto/field/" + f.Name, } assertNoError(polynomial.Generate(extInfo, filepath.Join(outputDir, "extensions", "polynomial"), i == 0, true, gen)) } diff --git a/internal/generator/polynomial/template/multilin.go.tmpl b/internal/generator/polynomial/template/multilin.go.tmpl index 50a08c10a5..adcd2b9e5c 100644 --- a/internal/generator/polynomial/template/multilin.go.tmpl +++ b/internal/generator/polynomial/template/multilin.go.tmpl @@ -1,5 +1,9 @@ import ( "{{.FieldPackagePath}}" +{{- if .ExtensionDegree}} + fr "{{.BaseFieldPackagePath}}" + basepoly "{{.BaseFieldPackagePath}}/polynomial" +{{- end}} "math/bits" "github.com/consensys/gnark-crypto/utils" ) @@ -49,6 +53,69 @@ func (m *MultiLin{{.ExtensionName}}) FoldParallel(r {{.ElementType}}) utils.Task } } +{{if .ExtensionDegree}} +// FoldFromBase sets m to the partial evaluation X₁=r of the multilinear polynomial whose +// hypercube evaluations are the base field elements b: +// +// m[i] = b[i] + r (b[i + len(b)/2] - b[i]) +// +// m's backing array is reused if it is large enough. +func (m *MultiLin{{.ExtensionName}}) FoldFromBase(b basepoly.MultiLin, r *{{.ElementType}}) { + mid := len(b) / 2 + m.resize(mid) + m.foldFromBase(b, r, 0, mid) +} + +// FoldFromBaseParallel is the parallel version of FoldFromBase. It sizes m, and returns +// a task that computes the entries m[start:end]. The task reads *r when it runs, so *r must +// not be modified until all tasks have completed. +func (m *MultiLin{{.ExtensionName}}) FoldFromBaseParallel(b basepoly.MultiLin, r *{{.ElementType}}) utils.Task { + mid := len(b) / 2 + m.resize(mid) + dst := *m + + return func(start, end int) { + dst.foldFromBase(b, r, start, end) + } +} + +// EvaluateBase evaluates, on the given coordinates, the multilinear polynomial whose hypercube +// evaluations are the base field elements b. It first folds b into m by the first coordinate +// (see FoldFromBase), and then folds m by the remaining coordinates. m serves as scratch space, +// is overwritten, and its backing array is reused if it is large enough. +// Unlike Evaluate, no copy of the table is made, and b is left untouched. +func (m *MultiLin{{.ExtensionName}}) EvaluateBase(b basepoly.MultiLin, coordinates []{{.ElementType}}) {{.ElementType}} { + if len(coordinates) == 0 { + m.resize(1) + return *(*m)[0].SetElement(&b[0]) + } + + m.FoldFromBase(b, &coordinates[0]) + for _, r := range coordinates[1:] { + m.Fold(r) + } + return (*m)[0] +} + +func (m *MultiLin{{.ExtensionName}}) resize(n int) { + if cap(*m) >= n { + *m = (*m)[:n] + } else { + *m = make(MultiLin{{.ExtensionName}}, n) + } +} + +func (m MultiLin{{.ExtensionName}}) foldFromBase(b basepoly.MultiLin, r *{{.ElementType}}, start, end int) { + mid := len(b) / 2 + var diff fr.Element + for i := start; i < end; i++ { + diff.Sub(&b[mid+i], &b[i]) + m[i].MulByElement(r, &diff) + m[i].AddElement(&m[i], &b[i]) + } +} + +{{end}} func (m MultiLin{{.ExtensionName}}) Sum() {{.ElementType}} { s := m[0] for i := 1; i < len(m); i++ { diff --git a/internal/generator/polynomial/template/multilin.test.go.tmpl b/internal/generator/polynomial/template/multilin.test.go.tmpl index 878b05aa90..d0521711da 100644 --- a/internal/generator/polynomial/template/multilin.test.go.tmpl +++ b/internal/generator/polynomial/template/multilin.test.go.tmpl @@ -1,4 +1,8 @@ import ( +{{- if .ExtensionDegree}} + fr "{{.BaseFieldPackagePath}}" + basepoly "{{.BaseFieldPackagePath}}/polynomial" +{{- end}} "{{.FieldPackagePath}}" "github.com/stretchr/testify/assert" "testing" @@ -79,3 +83,68 @@ func TestFoldedEqTable{{.ExtensionName}}(t *testing.T) { } } + +{{- if .ExtensionDegree}} + +func TestFoldFromBase{{.ExtensionName}}(t *testing.T) { + for _, n := range []int{2, 4, 8, 64} { + b := make(basepoly.MultiLin, n) + for i := range b { + b[i].MustSetRandom() + } + var r {{.ElementType}} + r.MustSetRandom() + + // reference: lift b to the extension and fold in place + lifted := make(MultiLin{{.ExtensionName}}, n) + for i := range b { + lifted[i].SetElement(&b[i]) + } + lifted.Fold(r) + + var got MultiLin{{.ExtensionName}} + got.FoldFromBase(b, &r) + assert.Equal(t, lifted, got) + + // parallel version, split in two tasks, reusing a dirty destination + par := make(MultiLin{{.ExtensionName}}, n) + task := par.FoldFromBaseParallel(b, &r) + mid := n / 2 + task(0, mid/2) + task(mid/2, mid) + assert.Equal(t, lifted, par) + } +} + +func TestEvaluateBase{{.ExtensionName}}(t *testing.T) { + for _, nbVars := range []int{0, 1, 2, 3, 6} { + b := make(basepoly.MultiLin, 1< Date: Mon, 5 Oct 2026 16:19:47 -0500 Subject: [PATCH 11/18] feat: subFromElement Signed-off-by: Arya Tabaie --- field/babybear/extensions/e2.go | 8 ++++++++ field/babybear/extensions/e2_test.go | 9 +++++++++ field/babybear/extensions/e4.go | 10 ++++++++++ field/babybear/extensions/e4_test.go | 9 +++++++++ field/babybear/extensions/e6.go | 12 ++++++++++++ field/babybear/extensions/e6_test.go | 9 +++++++++ field/goldilocks/extensions/e2.go | 8 ++++++++ field/goldilocks/extensions/e2_test.go | 9 +++++++++ field/koalabear/extensions/e2.go | 8 ++++++++ field/koalabear/extensions/e2_test.go | 9 +++++++++ field/koalabear/extensions/e4.go | 10 ++++++++++ field/koalabear/extensions/e4_test.go | 9 +++++++++ field/koalabear/extensions/e6.go | 12 ++++++++++++ field/koalabear/extensions/e6_test.go | 9 +++++++++ field/mamabear/extensions/e3.go | 9 +++++++++ field/mamabear/extensions/e3_test.go | 9 +++++++++ .../generator/field/template/extensions/e2.go.tmpl | 8 ++++++++ .../field/template/extensions/e2_test.go.tmpl | 9 +++++++++ .../generator/field/template/extensions/e3.go.tmpl | 9 +++++++++ .../field/template/extensions/e3_test.go.tmpl | 9 +++++++++ .../generator/field/template/extensions/e4.go.tmpl | 10 ++++++++++ .../field/template/extensions/e4_test.go.tmpl | 9 +++++++++ .../generator/field/template/extensions/e6.go.tmpl | 12 ++++++++++++ .../field/template/extensions/e6_test.go.tmpl | 9 +++++++++ 24 files changed, 224 insertions(+) diff --git a/field/babybear/extensions/e2.go b/field/babybear/extensions/e2.go index ae19692e74..f2f87cde71 100644 --- a/field/babybear/extensions/e2.go +++ b/field/babybear/extensions/e2.go @@ -105,6 +105,14 @@ func (z *E2) SubElement(x *E2, y *fr.Element) *E2 { return z } +// SubFromElement sets z to x - y, where x is an element of the base field embedded in E2, and returns z +func (z *E2) SubFromElement(x *fr.Element, y *E2) *E2 { + xc := *x + z.A0.Sub(&xc, &y.A0) + z.A1.Neg(&y.A1) + return z +} + // SetElement sets z to x, an element of the base field embedded in E2, and returns z func (z *E2) SetElement(x *fr.Element) *E2 { v := *x diff --git a/field/babybear/extensions/e2_test.go b/field/babybear/extensions/e2_test.go index a649b3c495..43e8647b61 100644 --- a/field/babybear/extensions/e2_test.go +++ b/field/babybear/extensions/e2_test.go @@ -654,6 +654,15 @@ func TestE2ElementOps(t *testing.T) { want.Sub(&x, &lifted) require.True(t, got.Equal(&want), "SubElement") + got.SubFromElement(&e, &x) + want.Sub(&lifted, &x) + require.True(t, got.Equal(&want), "SubFromElement") + + // aliasing of the receiver with the second operand + got.Set(&x) + got.SubFromElement(&e, &got) + require.True(t, got.Equal(&want), "SubFromElement alias") + // aliasing of the receiver with the first operand got.Set(&x) got.AddElement(&got, &e) diff --git a/field/babybear/extensions/e4.go b/field/babybear/extensions/e4.go index 3cab206a80..c506c716e4 100644 --- a/field/babybear/extensions/e4.go +++ b/field/babybear/extensions/e4.go @@ -120,6 +120,16 @@ func (z *E4) SubElement(x *E4, y *fr.Element) *E4 { return z } +// SubFromElement sets z to x - y, where x is an element of the base field embedded in E4, and returns z +func (z *E4) SubFromElement(x *fr.Element, y *E4) *E4 { + xc := *x + z.B0.A0.Sub(&xc, &y.B0.A0) + z.B0.A1.Neg(&y.B0.A1) + z.B1.A0.Neg(&y.B1.A0) + z.B1.A1.Neg(&y.B1.A1) + return z +} + // SetElement sets z to x, an element of the base field embedded in E4, and returns z func (z *E4) SetElement(x *fr.Element) *E4 { v := *x diff --git a/field/babybear/extensions/e4_test.go b/field/babybear/extensions/e4_test.go index 2c6189b849..12b2d9a97b 100644 --- a/field/babybear/extensions/e4_test.go +++ b/field/babybear/extensions/e4_test.go @@ -1230,6 +1230,15 @@ func TestE4ElementOps(t *testing.T) { want.Sub(&x, &lifted) require.True(t, got.Equal(&want), "SubElement") + got.SubFromElement(&e, &x) + want.Sub(&lifted, &x) + require.True(t, got.Equal(&want), "SubFromElement") + + // aliasing of the receiver with the second operand + got.Set(&x) + got.SubFromElement(&e, &got) + require.True(t, got.Equal(&want), "SubFromElement alias") + // aliasing of the receiver with the first operand got.Set(&x) got.AddElement(&got, &e) diff --git a/field/babybear/extensions/e6.go b/field/babybear/extensions/e6.go index f37f878415..4d0167ffd9 100644 --- a/field/babybear/extensions/e6.go +++ b/field/babybear/extensions/e6.go @@ -114,6 +114,18 @@ func (z *E6) SubElement(x *E6, y *fr.Element) *E6 { return z } +// SubFromElement sets z to x - y, where x is an element of the base field embedded in E6, and returns z +func (z *E6) SubFromElement(x *fr.Element, y *E6) *E6 { + xc := *x + z.B0.A0.Sub(&xc, &y.B0.A0) + z.B0.A1.Neg(&y.B0.A1) + z.B1.A0.Neg(&y.B1.A0) + z.B1.A1.Neg(&y.B1.A1) + z.B2.A0.Neg(&y.B2.A0) + z.B2.A1.Neg(&y.B2.A1) + return z +} + // SetElement sets z to x, an element of the base field embedded in E6, and returns z func (z *E6) SetElement(x *fr.Element) *E6 { v := *x diff --git a/field/babybear/extensions/e6_test.go b/field/babybear/extensions/e6_test.go index d13acbbf66..2f79477fa3 100644 --- a/field/babybear/extensions/e6_test.go +++ b/field/babybear/extensions/e6_test.go @@ -435,6 +435,15 @@ func TestE6ElementOps(t *testing.T) { want.Sub(&x, &lifted) require.True(t, got.Equal(&want), "SubElement") + got.SubFromElement(&e, &x) + want.Sub(&lifted, &x) + require.True(t, got.Equal(&want), "SubFromElement") + + // aliasing of the receiver with the second operand + got.Set(&x) + got.SubFromElement(&e, &got) + require.True(t, got.Equal(&want), "SubFromElement alias") + // aliasing of the receiver with the first operand got.Set(&x) got.AddElement(&got, &e) diff --git a/field/goldilocks/extensions/e2.go b/field/goldilocks/extensions/e2.go index fbffc7d3ff..0f86ede6f0 100644 --- a/field/goldilocks/extensions/e2.go +++ b/field/goldilocks/extensions/e2.go @@ -105,6 +105,14 @@ func (z *E2) SubElement(x *E2, y *fr.Element) *E2 { return z } +// SubFromElement sets z to x - y, where x is an element of the base field embedded in E2, and returns z +func (z *E2) SubFromElement(x *fr.Element, y *E2) *E2 { + xc := *x + z.A0.Sub(&xc, &y.A0) + z.A1.Neg(&y.A1) + return z +} + // SetElement sets z to x, an element of the base field embedded in E2, and returns z func (z *E2) SetElement(x *fr.Element) *E2 { v := *x diff --git a/field/goldilocks/extensions/e2_test.go b/field/goldilocks/extensions/e2_test.go index 730bc30a1c..d8b7b3ae6b 100644 --- a/field/goldilocks/extensions/e2_test.go +++ b/field/goldilocks/extensions/e2_test.go @@ -637,6 +637,15 @@ func TestE2ElementOps(t *testing.T) { want.Sub(&x, &lifted) require.True(t, got.Equal(&want), "SubElement") + got.SubFromElement(&e, &x) + want.Sub(&lifted, &x) + require.True(t, got.Equal(&want), "SubFromElement") + + // aliasing of the receiver with the second operand + got.Set(&x) + got.SubFromElement(&e, &got) + require.True(t, got.Equal(&want), "SubFromElement alias") + // aliasing of the receiver with the first operand got.Set(&x) got.AddElement(&got, &e) diff --git a/field/koalabear/extensions/e2.go b/field/koalabear/extensions/e2.go index 964cfbccee..3cb53d87fb 100644 --- a/field/koalabear/extensions/e2.go +++ b/field/koalabear/extensions/e2.go @@ -105,6 +105,14 @@ func (z *E2) SubElement(x *E2, y *fr.Element) *E2 { return z } +// SubFromElement sets z to x - y, where x is an element of the base field embedded in E2, and returns z +func (z *E2) SubFromElement(x *fr.Element, y *E2) *E2 { + xc := *x + z.A0.Sub(&xc, &y.A0) + z.A1.Neg(&y.A1) + return z +} + // SetElement sets z to x, an element of the base field embedded in E2, and returns z func (z *E2) SetElement(x *fr.Element) *E2 { v := *x diff --git a/field/koalabear/extensions/e2_test.go b/field/koalabear/extensions/e2_test.go index 4dea016c51..f1a6160839 100644 --- a/field/koalabear/extensions/e2_test.go +++ b/field/koalabear/extensions/e2_test.go @@ -654,6 +654,15 @@ func TestE2ElementOps(t *testing.T) { want.Sub(&x, &lifted) require.True(t, got.Equal(&want), "SubElement") + got.SubFromElement(&e, &x) + want.Sub(&lifted, &x) + require.True(t, got.Equal(&want), "SubFromElement") + + // aliasing of the receiver with the second operand + got.Set(&x) + got.SubFromElement(&e, &got) + require.True(t, got.Equal(&want), "SubFromElement alias") + // aliasing of the receiver with the first operand got.Set(&x) got.AddElement(&got, &e) diff --git a/field/koalabear/extensions/e4.go b/field/koalabear/extensions/e4.go index 2ea799059d..6163f06d72 100644 --- a/field/koalabear/extensions/e4.go +++ b/field/koalabear/extensions/e4.go @@ -120,6 +120,16 @@ func (z *E4) SubElement(x *E4, y *fr.Element) *E4 { return z } +// SubFromElement sets z to x - y, where x is an element of the base field embedded in E4, and returns z +func (z *E4) SubFromElement(x *fr.Element, y *E4) *E4 { + xc := *x + z.B0.A0.Sub(&xc, &y.B0.A0) + z.B0.A1.Neg(&y.B0.A1) + z.B1.A0.Neg(&y.B1.A0) + z.B1.A1.Neg(&y.B1.A1) + return z +} + // SetElement sets z to x, an element of the base field embedded in E4, and returns z func (z *E4) SetElement(x *fr.Element) *E4 { v := *x diff --git a/field/koalabear/extensions/e4_test.go b/field/koalabear/extensions/e4_test.go index 4ceea99856..1dd7db51cd 100644 --- a/field/koalabear/extensions/e4_test.go +++ b/field/koalabear/extensions/e4_test.go @@ -1230,6 +1230,15 @@ func TestE4ElementOps(t *testing.T) { want.Sub(&x, &lifted) require.True(t, got.Equal(&want), "SubElement") + got.SubFromElement(&e, &x) + want.Sub(&lifted, &x) + require.True(t, got.Equal(&want), "SubFromElement") + + // aliasing of the receiver with the second operand + got.Set(&x) + got.SubFromElement(&e, &got) + require.True(t, got.Equal(&want), "SubFromElement alias") + // aliasing of the receiver with the first operand got.Set(&x) got.AddElement(&got, &e) diff --git a/field/koalabear/extensions/e6.go b/field/koalabear/extensions/e6.go index 8d3054a7a9..342fb3e4e4 100644 --- a/field/koalabear/extensions/e6.go +++ b/field/koalabear/extensions/e6.go @@ -115,6 +115,18 @@ func (z *E6) SubElement(x *E6, y *fr.Element) *E6 { return z } +// SubFromElement sets z to x - y, where x is an element of the base field embedded in E6, and returns z +func (z *E6) SubFromElement(x *fr.Element, y *E6) *E6 { + xc := *x + z.B0.A0.Sub(&xc, &y.B0.A0) + z.B0.A1.Neg(&y.B0.A1) + z.B1.A0.Neg(&y.B1.A0) + z.B1.A1.Neg(&y.B1.A1) + z.B2.A0.Neg(&y.B2.A0) + z.B2.A1.Neg(&y.B2.A1) + return z +} + // SetElement sets z to x, an element of the base field embedded in E6, and returns z func (z *E6) SetElement(x *fr.Element) *E6 { v := *x diff --git a/field/koalabear/extensions/e6_test.go b/field/koalabear/extensions/e6_test.go index b14ada11f1..151fdfb846 100644 --- a/field/koalabear/extensions/e6_test.go +++ b/field/koalabear/extensions/e6_test.go @@ -435,6 +435,15 @@ func TestE6ElementOps(t *testing.T) { want.Sub(&x, &lifted) require.True(t, got.Equal(&want), "SubElement") + got.SubFromElement(&e, &x) + want.Sub(&lifted, &x) + require.True(t, got.Equal(&want), "SubFromElement") + + // aliasing of the receiver with the second operand + got.Set(&x) + got.SubFromElement(&e, &got) + require.True(t, got.Equal(&want), "SubFromElement alias") + // aliasing of the receiver with the first operand got.Set(&x) got.AddElement(&got, &e) diff --git a/field/mamabear/extensions/e3.go b/field/mamabear/extensions/e3.go index 748929b8f9..ac6343051f 100644 --- a/field/mamabear/extensions/e3.go +++ b/field/mamabear/extensions/e3.go @@ -82,6 +82,15 @@ func (z *E3) SubElement(x *E3, y *fr.Element) *E3 { return z } +// SubFromElement sets z to x - y, where x is an element of the base field embedded in E3, and returns z +func (z *E3) SubFromElement(x *fr.Element, y *E3) *E3 { + xc := *x + z.A0.Sub(&xc, &y.A0) + z.A1.Neg(&y.A1) + z.A2.Neg(&y.A2) + return z +} + // SetElement sets z to x, an element of the base field embedded in E3, and returns z func (z *E3) SetElement(x *fr.Element) *E3 { v := *x diff --git a/field/mamabear/extensions/e3_test.go b/field/mamabear/extensions/e3_test.go index 590b277321..0bbf06a2b0 100644 --- a/field/mamabear/extensions/e3_test.go +++ b/field/mamabear/extensions/e3_test.go @@ -566,6 +566,15 @@ func TestE3ElementOps(t *testing.T) { want.Sub(&x, &lifted) require.True(t, got.Equal(&want), "SubElement") + got.SubFromElement(&e, &x) + want.Sub(&lifted, &x) + require.True(t, got.Equal(&want), "SubFromElement") + + // aliasing of the receiver with the second operand + got.Set(&x) + got.SubFromElement(&e, &got) + require.True(t, got.Equal(&want), "SubFromElement alias") + // aliasing of the receiver with the first operand got.Set(&x) got.AddElement(&got, &e) diff --git a/internal/generator/field/template/extensions/e2.go.tmpl b/internal/generator/field/template/extensions/e2.go.tmpl index 3032902239..dc50b44ad0 100644 --- a/internal/generator/field/template/extensions/e2.go.tmpl +++ b/internal/generator/field/template/extensions/e2.go.tmpl @@ -97,6 +97,14 @@ func (z *E2) SubElement(x *E2, y *fr.Element) *E2 { return z } +// SubFromElement sets z to x - y, where x is an element of the base field embedded in E2, and returns z +func (z *E2) SubFromElement(x *fr.Element, y *E2) *E2 { + xc := *x + z.A0.Sub(&xc, &y.A0) + z.A1.Neg(&y.A1) + return z +} + // SetElement sets z to x, an element of the base field embedded in E2, and returns z func (z *E2) SetElement(x *fr.Element) *E2 { v := *x diff --git a/internal/generator/field/template/extensions/e2_test.go.tmpl b/internal/generator/field/template/extensions/e2_test.go.tmpl index ef6a39367d..171b017b1f 100644 --- a/internal/generator/field/template/extensions/e2_test.go.tmpl +++ b/internal/generator/field/template/extensions/e2_test.go.tmpl @@ -658,6 +658,15 @@ func TestE2ElementOps(t *testing.T) { want.Sub(&x, &lifted) require.True(t, got.Equal(&want), "SubElement") + got.SubFromElement(&e, &x) + want.Sub(&lifted, &x) + require.True(t, got.Equal(&want), "SubFromElement") + + // aliasing of the receiver with the second operand + got.Set(&x) + got.SubFromElement(&e, &got) + require.True(t, got.Equal(&want), "SubFromElement alias") + // aliasing of the receiver with the first operand got.Set(&x) got.AddElement(&got, &e) diff --git a/internal/generator/field/template/extensions/e3.go.tmpl b/internal/generator/field/template/extensions/e3.go.tmpl index b5f9e31ed8..9f16f0fb41 100644 --- a/internal/generator/field/template/extensions/e3.go.tmpl +++ b/internal/generator/field/template/extensions/e3.go.tmpl @@ -74,6 +74,15 @@ func (z *E3) SubElement(x *E3, y *fr.Element) *E3 { return z } +// SubFromElement sets z to x - y, where x is an element of the base field embedded in E3, and returns z +func (z *E3) SubFromElement(x *fr.Element, y *E3) *E3 { + xc := *x + z.A0.Sub(&xc, &y.A0) + z.A1.Neg(&y.A1) + z.A2.Neg(&y.A2) + return z +} + // SetElement sets z to x, an element of the base field embedded in E3, and returns z func (z *E3) SetElement(x *fr.Element) *E3 { v := *x diff --git a/internal/generator/field/template/extensions/e3_test.go.tmpl b/internal/generator/field/template/extensions/e3_test.go.tmpl index 72c9878191..b4944ab58a 100644 --- a/internal/generator/field/template/extensions/e3_test.go.tmpl +++ b/internal/generator/field/template/extensions/e3_test.go.tmpl @@ -557,6 +557,15 @@ func TestE3ElementOps(t *testing.T) { want.Sub(&x, &lifted) require.True(t, got.Equal(&want), "SubElement") + got.SubFromElement(&e, &x) + want.Sub(&lifted, &x) + require.True(t, got.Equal(&want), "SubFromElement") + + // aliasing of the receiver with the second operand + got.Set(&x) + got.SubFromElement(&e, &got) + require.True(t, got.Equal(&want), "SubFromElement alias") + // aliasing of the receiver with the first operand got.Set(&x) got.AddElement(&got, &e) diff --git a/internal/generator/field/template/extensions/e4.go.tmpl b/internal/generator/field/template/extensions/e4.go.tmpl index dd7fd2633a..62ab7dbb5d 100644 --- a/internal/generator/field/template/extensions/e4.go.tmpl +++ b/internal/generator/field/template/extensions/e4.go.tmpl @@ -112,6 +112,16 @@ func (z *E4) SubElement(x *E4, y *fr.Element) *E4 { return z } +// SubFromElement sets z to x - y, where x is an element of the base field embedded in E4, and returns z +func (z *E4) SubFromElement(x *fr.Element, y *E4) *E4 { + xc := *x + z.B0.A0.Sub(&xc, &y.B0.A0) + z.B0.A1.Neg(&y.B0.A1) + z.B1.A0.Neg(&y.B1.A0) + z.B1.A1.Neg(&y.B1.A1) + return z +} + // SetElement sets z to x, an element of the base field embedded in E4, and returns z func (z *E4) SetElement(x *fr.Element) *E4 { v := *x diff --git a/internal/generator/field/template/extensions/e4_test.go.tmpl b/internal/generator/field/template/extensions/e4_test.go.tmpl index 174c0702e9..42f099c10a 100644 --- a/internal/generator/field/template/extensions/e4_test.go.tmpl +++ b/internal/generator/field/template/extensions/e4_test.go.tmpl @@ -1228,6 +1228,15 @@ func TestE4ElementOps(t *testing.T) { want.Sub(&x, &lifted) require.True(t, got.Equal(&want), "SubElement") + got.SubFromElement(&e, &x) + want.Sub(&lifted, &x) + require.True(t, got.Equal(&want), "SubFromElement") + + // aliasing of the receiver with the second operand + got.Set(&x) + got.SubFromElement(&e, &got) + require.True(t, got.Equal(&want), "SubFromElement alias") + // aliasing of the receiver with the first operand got.Set(&x) got.AddElement(&got, &e) diff --git a/internal/generator/field/template/extensions/e6.go.tmpl b/internal/generator/field/template/extensions/e6.go.tmpl index 8a87423a23..07ac75b1b8 100644 --- a/internal/generator/field/template/extensions/e6.go.tmpl +++ b/internal/generator/field/template/extensions/e6.go.tmpl @@ -109,6 +109,18 @@ func (z *E6) SubElement(x *E6, y *fr.Element) *E6 { return z } +// SubFromElement sets z to x - y, where x is an element of the base field embedded in E6, and returns z +func (z *E6) SubFromElement(x *fr.Element, y *E6) *E6 { + xc := *x + z.B0.A0.Sub(&xc, &y.B0.A0) + z.B0.A1.Neg(&y.B0.A1) + z.B1.A0.Neg(&y.B1.A0) + z.B1.A1.Neg(&y.B1.A1) + z.B2.A0.Neg(&y.B2.A0) + z.B2.A1.Neg(&y.B2.A1) + return z +} + // SetElement sets z to x, an element of the base field embedded in E6, and returns z func (z *E6) SetElement(x *fr.Element) *E6 { v := *x diff --git a/internal/generator/field/template/extensions/e6_test.go.tmpl b/internal/generator/field/template/extensions/e6_test.go.tmpl index e7f0d8b35b..e6881e96ac 100644 --- a/internal/generator/field/template/extensions/e6_test.go.tmpl +++ b/internal/generator/field/template/extensions/e6_test.go.tmpl @@ -425,6 +425,15 @@ func TestE6ElementOps(t *testing.T) { want.Sub(&x, &lifted) require.True(t, got.Equal(&want), "SubElement") + got.SubFromElement(&e, &x) + want.Sub(&lifted, &x) + require.True(t, got.Equal(&want), "SubFromElement") + + // aliasing of the receiver with the second operand + got.Set(&x) + got.SubFromElement(&e, &got) + require.True(t, got.Equal(&want), "SubFromElement alias") + // aliasing of the receiver with the first operand got.Set(&x) got.AddElement(&got, &e) From 302a7debb8761353d42bf84e7153a9c66a1e8c24 Mon Sep 17 00:00:00 2001 From: Arya Tabaie Date: Tue, 6 Oct 2026 15:08:33 -0500 Subject: [PATCH 12/18] feat: export external and internal matrices Signed-off-by: Arya Tabaie --- ecc/bls12-377/fr/poseidon2/poseidon2.go | 38 +++++ ecc/bls12-377/fr/poseidon2/poseidon2_test.go | 46 +++++++ ecc/bls12-381/fr/poseidon2/poseidon2.go | 38 +++++ ecc/bls12-381/fr/poseidon2/poseidon2_test.go | 46 +++++++ ecc/bls24-315/fr/poseidon2/poseidon2.go | 38 +++++ ecc/bls24-315/fr/poseidon2/poseidon2_test.go | 46 +++++++ ecc/bls24-317/fr/poseidon2/poseidon2.go | 38 +++++ ecc/bls24-317/fr/poseidon2/poseidon2_test.go | 46 +++++++ ecc/bn254/fr/poseidon2/poseidon2.go | 41 ++++++ ecc/bn254/fr/poseidon2/poseidon2_test.go | 50 +++++++ ecc/bw6-633/fr/poseidon2/poseidon2.go | 38 +++++ ecc/bw6-633/fr/poseidon2/poseidon2_test.go | 46 +++++++ ecc/bw6-761/fr/poseidon2/poseidon2.go | 38 +++++ ecc/bw6-761/fr/poseidon2/poseidon2_test.go | 46 +++++++ ecc/grumpkin/fr/poseidon2/poseidon2.go | 38 +++++ ecc/grumpkin/fr/poseidon2/poseidon2_test.go | 46 +++++++ field/babybear/poseidon2/poseidon2.go | 31 +++++ field/babybear/poseidon2/poseidon2_test.go | 46 +++++++ field/goldilocks/poseidon2/poseidon2.go | 31 +++++ field/goldilocks/poseidon2/poseidon2_test.go | 46 +++++++ field/koalabear/poseidon2/poseidon2.go | 31 +++++ field/koalabear/poseidon2/poseidon2_test.go | 46 +++++++ field/mamabear/poseidon2/poseidon2.go | 31 +++++ field/mamabear/poseidon2/poseidon2_test.go | 46 +++++++ .../hash/poseidon2/template/poseidon2.go.tmpl | 45 ++++++ .../poseidon2/template/poseidon2.test.go.tmpl | 54 +++++++- .../template/poseidon2/poseidon2.go.tmpl | 31 +++++ .../template/poseidon2/poseidon2_test.go.tmpl | 47 +++++++ internal/poseidon2/matrices.go | 121 ++++++++++++++++ internal/poseidon2/matrices_test.go | 130 ++++++++++++++++++ 30 files changed, 1414 insertions(+), 1 deletion(-) create mode 100644 internal/poseidon2/matrices.go create mode 100644 internal/poseidon2/matrices_test.go diff --git a/ecc/bls12-377/fr/poseidon2/poseidon2.go b/ecc/bls12-377/fr/poseidon2/poseidon2.go index 73558f2e92..ce9c17f44a 100644 --- a/ecc/bls12-377/fr/poseidon2/poseidon2.go +++ b/ecc/bls12-377/fr/poseidon2/poseidon2.go @@ -12,6 +12,7 @@ import ( "golang.org/x/crypto/sha3" "github.com/consensys/gnark-crypto/ecc/bls12-377/fr" + internalposeidon2 "github.com/consensys/gnark-crypto/internal/poseidon2" ) var ( @@ -121,6 +122,43 @@ type Permutation struct { params *Parameters } +// ExternalMatrix returns the dense external matrix M_E of the permutation, which +// is applied to the state before the first round and after each full round. +// With I and J the identity and the all-ones matrices, and M4 the 4×4 block of +// the Poseidon2 paper (rows 5 7 1 3 / 4 6 1 1 / 1 3 5 7 / 1 1 4 6), M_E is +// +// - I + J for width 2 and 3; +// - M4 for width 4; +// - circ(2·M4, M4, ..., M4) = (I + J) ⊗ M4 for width 4k with k > 1. +func (p *Parameters) ExternalMatrix() [][]fr.Element { + return internalposeidon2.ExternalMatrix[fr.Element](p.Width, internalposeidon2.Paper) +} + +// InternalMatrix returns the dense internal matrix M_I of the permutation, which +// is applied to the state after each partial round. With J the all-ones matrix, +// +// M_I = J + diag(d), +// +// that is, for a state x, (M_I·x)_i = Σ_j x_j + d_i·x_i, where d is the diagonal +// of the permutation. +func (p *Parameters) InternalMatrix() [][]fr.Element { + var d []fr.Element + switch p.Width { + case 2: + d = make([]fr.Element, 2) + d[0].SetOne() + d[1].SetUint64(2) + case 3: + d = make([]fr.Element, 3) + d[0].SetOne() + d[1].SetOne() + d[2].SetUint64(2) + default: + panic("only Width=2,3 are supported") + } + return internalposeidon2.InternalMatrix(d) +} + // NewPermutation returns a new Poseidon2 permutation instance. func NewPermutation(t, rf, rp int) *Permutation { if t < 2 || t > 3 { diff --git a/ecc/bls12-377/fr/poseidon2/poseidon2_test.go b/ecc/bls12-377/fr/poseidon2/poseidon2_test.go index 28647f8d12..17b7c0b70b 100644 --- a/ecc/bls12-377/fr/poseidon2/poseidon2_test.go +++ b/ecc/bls12-377/fr/poseidon2/poseidon2_test.go @@ -109,3 +109,49 @@ func TestHashReset(t *testing.T) { require.Equal(t, res, h.Sum(nil)) } + +// denseMatrixMul returns m·x. +func denseMatrixMul(m [][]fr.Element, x []fr.Element) []fr.Element { + res := make([]fr.Element, len(m)) + for i := range m { + var tmp fr.Element + for j := range m[i] { + tmp.Mul(&m[i][j], &x[j]) + res[i].Add(&res[i], &tmp) + } + } + return res +} + +// TestDenseMatricesMatchInPlace checks that the dense matrices act on a random +// state as the in-place multiplications of the permutation. +func TestDenseMatricesMatchInPlace(t *testing.T) { + for _, tc := range []struct { + width, nbFullRounds, nbPartialRounds int + }{ + {2, 8, 56}, + {3, 8, 56}, + } { + h := NewPermutation(tc.width, tc.nbFullRounds, tc.nbPartialRounds) + + for name, step := range map[string]struct { + dense [][]fr.Element + inPlace func([]fr.Element) + }{ + "external": {h.params.ExternalMatrix(), h.matMulExternalInPlace}, + "internal": {h.params.InternalMatrix(), h.matMulInternalInPlace}, + } { + x := make([]fr.Element, tc.width) + for i := range x { + x[i].MustSetRandom() + } + want := denseMatrixMul(step.dense, x) + step.inPlace(x) + for i := range x { + if !x[i].Equal(&want[i]) { + t.Fatalf("width %d, %s matrix: dense and in-place multiplications differ at index %d", tc.width, name, i) + } + } + } + } +} diff --git a/ecc/bls12-381/fr/poseidon2/poseidon2.go b/ecc/bls12-381/fr/poseidon2/poseidon2.go index b98a7b4846..70d662a45f 100644 --- a/ecc/bls12-381/fr/poseidon2/poseidon2.go +++ b/ecc/bls12-381/fr/poseidon2/poseidon2.go @@ -12,6 +12,7 @@ import ( "golang.org/x/crypto/sha3" "github.com/consensys/gnark-crypto/ecc/bls12-381/fr" + internalposeidon2 "github.com/consensys/gnark-crypto/internal/poseidon2" ) var ( @@ -121,6 +122,43 @@ type Permutation struct { params *Parameters } +// ExternalMatrix returns the dense external matrix M_E of the permutation, which +// is applied to the state before the first round and after each full round. +// With I and J the identity and the all-ones matrices, and M4 the 4×4 block of +// the Poseidon2 paper (rows 5 7 1 3 / 4 6 1 1 / 1 3 5 7 / 1 1 4 6), M_E is +// +// - I + J for width 2 and 3; +// - M4 for width 4; +// - circ(2·M4, M4, ..., M4) = (I + J) ⊗ M4 for width 4k with k > 1. +func (p *Parameters) ExternalMatrix() [][]fr.Element { + return internalposeidon2.ExternalMatrix[fr.Element](p.Width, internalposeidon2.Paper) +} + +// InternalMatrix returns the dense internal matrix M_I of the permutation, which +// is applied to the state after each partial round. With J the all-ones matrix, +// +// M_I = J + diag(d), +// +// that is, for a state x, (M_I·x)_i = Σ_j x_j + d_i·x_i, where d is the diagonal +// of the permutation. +func (p *Parameters) InternalMatrix() [][]fr.Element { + var d []fr.Element + switch p.Width { + case 2: + d = make([]fr.Element, 2) + d[0].SetOne() + d[1].SetUint64(2) + case 3: + d = make([]fr.Element, 3) + d[0].SetOne() + d[1].SetOne() + d[2].SetUint64(2) + default: + panic("only Width=2,3 are supported") + } + return internalposeidon2.InternalMatrix(d) +} + // NewPermutation returns a new Poseidon2 permutation instance. func NewPermutation(t, rf, rp int) *Permutation { if t < 2 || t > 3 { diff --git a/ecc/bls12-381/fr/poseidon2/poseidon2_test.go b/ecc/bls12-381/fr/poseidon2/poseidon2_test.go index c577df2f7c..324edceef5 100644 --- a/ecc/bls12-381/fr/poseidon2/poseidon2_test.go +++ b/ecc/bls12-381/fr/poseidon2/poseidon2_test.go @@ -109,3 +109,49 @@ func TestHashReset(t *testing.T) { require.Equal(t, res, h.Sum(nil)) } + +// denseMatrixMul returns m·x. +func denseMatrixMul(m [][]fr.Element, x []fr.Element) []fr.Element { + res := make([]fr.Element, len(m)) + for i := range m { + var tmp fr.Element + for j := range m[i] { + tmp.Mul(&m[i][j], &x[j]) + res[i].Add(&res[i], &tmp) + } + } + return res +} + +// TestDenseMatricesMatchInPlace checks that the dense matrices act on a random +// state as the in-place multiplications of the permutation. +func TestDenseMatricesMatchInPlace(t *testing.T) { + for _, tc := range []struct { + width, nbFullRounds, nbPartialRounds int + }{ + {2, 8, 56}, + {3, 8, 56}, + } { + h := NewPermutation(tc.width, tc.nbFullRounds, tc.nbPartialRounds) + + for name, step := range map[string]struct { + dense [][]fr.Element + inPlace func([]fr.Element) + }{ + "external": {h.params.ExternalMatrix(), h.matMulExternalInPlace}, + "internal": {h.params.InternalMatrix(), h.matMulInternalInPlace}, + } { + x := make([]fr.Element, tc.width) + for i := range x { + x[i].MustSetRandom() + } + want := denseMatrixMul(step.dense, x) + step.inPlace(x) + for i := range x { + if !x[i].Equal(&want[i]) { + t.Fatalf("width %d, %s matrix: dense and in-place multiplications differ at index %d", tc.width, name, i) + } + } + } + } +} diff --git a/ecc/bls24-315/fr/poseidon2/poseidon2.go b/ecc/bls24-315/fr/poseidon2/poseidon2.go index 6c4be772a2..6d233573c3 100644 --- a/ecc/bls24-315/fr/poseidon2/poseidon2.go +++ b/ecc/bls24-315/fr/poseidon2/poseidon2.go @@ -12,6 +12,7 @@ import ( "golang.org/x/crypto/sha3" "github.com/consensys/gnark-crypto/ecc/bls24-315/fr" + internalposeidon2 "github.com/consensys/gnark-crypto/internal/poseidon2" ) var ( @@ -121,6 +122,43 @@ type Permutation struct { params *Parameters } +// ExternalMatrix returns the dense external matrix M_E of the permutation, which +// is applied to the state before the first round and after each full round. +// With I and J the identity and the all-ones matrices, and M4 the 4×4 block of +// the Poseidon2 paper (rows 5 7 1 3 / 4 6 1 1 / 1 3 5 7 / 1 1 4 6), M_E is +// +// - I + J for width 2 and 3; +// - M4 for width 4; +// - circ(2·M4, M4, ..., M4) = (I + J) ⊗ M4 for width 4k with k > 1. +func (p *Parameters) ExternalMatrix() [][]fr.Element { + return internalposeidon2.ExternalMatrix[fr.Element](p.Width, internalposeidon2.Paper) +} + +// InternalMatrix returns the dense internal matrix M_I of the permutation, which +// is applied to the state after each partial round. With J the all-ones matrix, +// +// M_I = J + diag(d), +// +// that is, for a state x, (M_I·x)_i = Σ_j x_j + d_i·x_i, where d is the diagonal +// of the permutation. +func (p *Parameters) InternalMatrix() [][]fr.Element { + var d []fr.Element + switch p.Width { + case 2: + d = make([]fr.Element, 2) + d[0].SetOne() + d[1].SetUint64(2) + case 3: + d = make([]fr.Element, 3) + d[0].SetOne() + d[1].SetOne() + d[2].SetUint64(2) + default: + panic("only Width=2,3 are supported") + } + return internalposeidon2.InternalMatrix(d) +} + // NewPermutation returns a new Poseidon2 permutation instance. func NewPermutation(t, rf, rp int) *Permutation { if t < 2 || t > 3 { diff --git a/ecc/bls24-315/fr/poseidon2/poseidon2_test.go b/ecc/bls24-315/fr/poseidon2/poseidon2_test.go index 1a2f1daedc..e35d355476 100644 --- a/ecc/bls24-315/fr/poseidon2/poseidon2_test.go +++ b/ecc/bls24-315/fr/poseidon2/poseidon2_test.go @@ -109,3 +109,49 @@ func TestHashReset(t *testing.T) { require.Equal(t, res, h.Sum(nil)) } + +// denseMatrixMul returns m·x. +func denseMatrixMul(m [][]fr.Element, x []fr.Element) []fr.Element { + res := make([]fr.Element, len(m)) + for i := range m { + var tmp fr.Element + for j := range m[i] { + tmp.Mul(&m[i][j], &x[j]) + res[i].Add(&res[i], &tmp) + } + } + return res +} + +// TestDenseMatricesMatchInPlace checks that the dense matrices act on a random +// state as the in-place multiplications of the permutation. +func TestDenseMatricesMatchInPlace(t *testing.T) { + for _, tc := range []struct { + width, nbFullRounds, nbPartialRounds int + }{ + {2, 8, 56}, + {3, 8, 56}, + } { + h := NewPermutation(tc.width, tc.nbFullRounds, tc.nbPartialRounds) + + for name, step := range map[string]struct { + dense [][]fr.Element + inPlace func([]fr.Element) + }{ + "external": {h.params.ExternalMatrix(), h.matMulExternalInPlace}, + "internal": {h.params.InternalMatrix(), h.matMulInternalInPlace}, + } { + x := make([]fr.Element, tc.width) + for i := range x { + x[i].MustSetRandom() + } + want := denseMatrixMul(step.dense, x) + step.inPlace(x) + for i := range x { + if !x[i].Equal(&want[i]) { + t.Fatalf("width %d, %s matrix: dense and in-place multiplications differ at index %d", tc.width, name, i) + } + } + } + } +} diff --git a/ecc/bls24-317/fr/poseidon2/poseidon2.go b/ecc/bls24-317/fr/poseidon2/poseidon2.go index 7e1061a638..e1809d2f86 100644 --- a/ecc/bls24-317/fr/poseidon2/poseidon2.go +++ b/ecc/bls24-317/fr/poseidon2/poseidon2.go @@ -12,6 +12,7 @@ import ( "golang.org/x/crypto/sha3" "github.com/consensys/gnark-crypto/ecc/bls24-317/fr" + internalposeidon2 "github.com/consensys/gnark-crypto/internal/poseidon2" ) var ( @@ -121,6 +122,43 @@ type Permutation struct { params *Parameters } +// ExternalMatrix returns the dense external matrix M_E of the permutation, which +// is applied to the state before the first round and after each full round. +// With I and J the identity and the all-ones matrices, and M4 the 4×4 block of +// the Poseidon2 paper (rows 5 7 1 3 / 4 6 1 1 / 1 3 5 7 / 1 1 4 6), M_E is +// +// - I + J for width 2 and 3; +// - M4 for width 4; +// - circ(2·M4, M4, ..., M4) = (I + J) ⊗ M4 for width 4k with k > 1. +func (p *Parameters) ExternalMatrix() [][]fr.Element { + return internalposeidon2.ExternalMatrix[fr.Element](p.Width, internalposeidon2.Paper) +} + +// InternalMatrix returns the dense internal matrix M_I of the permutation, which +// is applied to the state after each partial round. With J the all-ones matrix, +// +// M_I = J + diag(d), +// +// that is, for a state x, (M_I·x)_i = Σ_j x_j + d_i·x_i, where d is the diagonal +// of the permutation. +func (p *Parameters) InternalMatrix() [][]fr.Element { + var d []fr.Element + switch p.Width { + case 2: + d = make([]fr.Element, 2) + d[0].SetOne() + d[1].SetUint64(2) + case 3: + d = make([]fr.Element, 3) + d[0].SetOne() + d[1].SetOne() + d[2].SetUint64(2) + default: + panic("only Width=2,3 are supported") + } + return internalposeidon2.InternalMatrix(d) +} + // NewPermutation returns a new Poseidon2 permutation instance. func NewPermutation(t, rf, rp int) *Permutation { if t < 2 || t > 3 { diff --git a/ecc/bls24-317/fr/poseidon2/poseidon2_test.go b/ecc/bls24-317/fr/poseidon2/poseidon2_test.go index b89e8741f5..1165a8ca8b 100644 --- a/ecc/bls24-317/fr/poseidon2/poseidon2_test.go +++ b/ecc/bls24-317/fr/poseidon2/poseidon2_test.go @@ -109,3 +109,49 @@ func TestHashReset(t *testing.T) { require.Equal(t, res, h.Sum(nil)) } + +// denseMatrixMul returns m·x. +func denseMatrixMul(m [][]fr.Element, x []fr.Element) []fr.Element { + res := make([]fr.Element, len(m)) + for i := range m { + var tmp fr.Element + for j := range m[i] { + tmp.Mul(&m[i][j], &x[j]) + res[i].Add(&res[i], &tmp) + } + } + return res +} + +// TestDenseMatricesMatchInPlace checks that the dense matrices act on a random +// state as the in-place multiplications of the permutation. +func TestDenseMatricesMatchInPlace(t *testing.T) { + for _, tc := range []struct { + width, nbFullRounds, nbPartialRounds int + }{ + {2, 8, 56}, + {3, 8, 56}, + } { + h := NewPermutation(tc.width, tc.nbFullRounds, tc.nbPartialRounds) + + for name, step := range map[string]struct { + dense [][]fr.Element + inPlace func([]fr.Element) + }{ + "external": {h.params.ExternalMatrix(), h.matMulExternalInPlace}, + "internal": {h.params.InternalMatrix(), h.matMulInternalInPlace}, + } { + x := make([]fr.Element, tc.width) + for i := range x { + x[i].MustSetRandom() + } + want := denseMatrixMul(step.dense, x) + step.inPlace(x) + for i := range x { + if !x[i].Equal(&want[i]) { + t.Fatalf("width %d, %s matrix: dense and in-place multiplications differ at index %d", tc.width, name, i) + } + } + } + } +} diff --git a/ecc/bn254/fr/poseidon2/poseidon2.go b/ecc/bn254/fr/poseidon2/poseidon2.go index 66ca9e8830..09ab10e984 100644 --- a/ecc/bn254/fr/poseidon2/poseidon2.go +++ b/ecc/bn254/fr/poseidon2/poseidon2.go @@ -12,6 +12,7 @@ import ( "golang.org/x/crypto/sha3" "github.com/consensys/gnark-crypto/ecc/bn254/fr" + internalposeidon2 "github.com/consensys/gnark-crypto/internal/poseidon2" ) var ( @@ -140,6 +141,46 @@ type Permutation struct { params *Parameters } +// ExternalMatrix returns the dense external matrix M_E of the permutation, which +// is applied to the state before the first round and after each full round. +// With I and J the identity and the all-ones matrices, and M4 the 4×4 block of +// the Poseidon2 paper (rows 5 7 1 3 / 4 6 1 1 / 1 3 5 7 / 1 1 4 6), M_E is +// +// - I + J for width 2 and 3; +// - M4 for width 4; +// - circ(2·M4, M4, ..., M4) = (I + J) ⊗ M4 for width 4k with k > 1. +func (p *Parameters) ExternalMatrix() [][]fr.Element { + return internalposeidon2.ExternalMatrix[fr.Element](p.Width, internalposeidon2.Paper) +} + +// InternalMatrix returns the dense internal matrix M_I of the permutation, which +// is applied to the state after each partial round. With J the all-ones matrix, +// +// M_I = J + diag(d), +// +// that is, for a state x, (M_I·x)_i = Σ_j x_j + d_i·x_i, where d is the diagonal +// of the permutation (DiagM1 for the widths of 4 and more). +func (p *Parameters) InternalMatrix() [][]fr.Element { + var d []fr.Element + switch p.Width { + case 2: + d = make([]fr.Element, 2) + d[0].SetOne() + d[1].SetUint64(2) + case 3: + d = make([]fr.Element, 3) + d[0].SetOne() + d[1].SetOne() + d[2].SetUint64(2) + default: + if len(p.DiagM1) != p.Width { + panic("missing internal matrix diagonal") + } + d = p.DiagM1 + } + return internalposeidon2.InternalMatrix(d) +} + // NewPermutation returns a new Poseidon2 permutation instance. func NewPermutation(t, rf, rp int) *Permutation { if t == 2 || t == 3 { diff --git a/ecc/bn254/fr/poseidon2/poseidon2_test.go b/ecc/bn254/fr/poseidon2/poseidon2_test.go index 9990070916..e672ca94fa 100644 --- a/ecc/bn254/fr/poseidon2/poseidon2_test.go +++ b/ecc/bn254/fr/poseidon2/poseidon2_test.go @@ -193,3 +193,53 @@ func TestHashReset(t *testing.T) { require.Equal(t, res, h.Sum(nil)) } + +// denseMatrixMul returns m·x. +func denseMatrixMul(m [][]fr.Element, x []fr.Element) []fr.Element { + res := make([]fr.Element, len(m)) + for i := range m { + var tmp fr.Element + for j := range m[i] { + tmp.Mul(&m[i][j], &x[j]) + res[i].Add(&res[i], &tmp) + } + } + return res +} + +// TestDenseMatricesMatchInPlace checks that the dense matrices act on a random +// state as the in-place multiplications of the permutation. +func TestDenseMatricesMatchInPlace(t *testing.T) { + for _, tc := range []struct { + width, nbFullRounds, nbPartialRounds int + }{ + {2, 8, 56}, + {3, 8, 56}, + {4, 8, 56}, + {8, 8, 57}, + {12, 8, 57}, + {16, 8, 57}, + } { + h := NewPermutation(tc.width, tc.nbFullRounds, tc.nbPartialRounds) + + for name, step := range map[string]struct { + dense [][]fr.Element + inPlace func([]fr.Element) + }{ + "external": {h.params.ExternalMatrix(), h.matMulExternalInPlace}, + "internal": {h.params.InternalMatrix(), h.matMulInternalInPlace}, + } { + x := make([]fr.Element, tc.width) + for i := range x { + x[i].MustSetRandom() + } + want := denseMatrixMul(step.dense, x) + step.inPlace(x) + for i := range x { + if !x[i].Equal(&want[i]) { + t.Fatalf("width %d, %s matrix: dense and in-place multiplications differ at index %d", tc.width, name, i) + } + } + } + } +} diff --git a/ecc/bw6-633/fr/poseidon2/poseidon2.go b/ecc/bw6-633/fr/poseidon2/poseidon2.go index e577b9dd88..35b47b266f 100644 --- a/ecc/bw6-633/fr/poseidon2/poseidon2.go +++ b/ecc/bw6-633/fr/poseidon2/poseidon2.go @@ -12,6 +12,7 @@ import ( "golang.org/x/crypto/sha3" "github.com/consensys/gnark-crypto/ecc/bw6-633/fr" + internalposeidon2 "github.com/consensys/gnark-crypto/internal/poseidon2" ) var ( @@ -121,6 +122,43 @@ type Permutation struct { params *Parameters } +// ExternalMatrix returns the dense external matrix M_E of the permutation, which +// is applied to the state before the first round and after each full round. +// With I and J the identity and the all-ones matrices, and M4 the 4×4 block of +// the Poseidon2 paper (rows 5 7 1 3 / 4 6 1 1 / 1 3 5 7 / 1 1 4 6), M_E is +// +// - I + J for width 2 and 3; +// - M4 for width 4; +// - circ(2·M4, M4, ..., M4) = (I + J) ⊗ M4 for width 4k with k > 1. +func (p *Parameters) ExternalMatrix() [][]fr.Element { + return internalposeidon2.ExternalMatrix[fr.Element](p.Width, internalposeidon2.Paper) +} + +// InternalMatrix returns the dense internal matrix M_I of the permutation, which +// is applied to the state after each partial round. With J the all-ones matrix, +// +// M_I = J + diag(d), +// +// that is, for a state x, (M_I·x)_i = Σ_j x_j + d_i·x_i, where d is the diagonal +// of the permutation. +func (p *Parameters) InternalMatrix() [][]fr.Element { + var d []fr.Element + switch p.Width { + case 2: + d = make([]fr.Element, 2) + d[0].SetOne() + d[1].SetUint64(2) + case 3: + d = make([]fr.Element, 3) + d[0].SetOne() + d[1].SetOne() + d[2].SetUint64(2) + default: + panic("only Width=2,3 are supported") + } + return internalposeidon2.InternalMatrix(d) +} + // NewPermutation returns a new Poseidon2 permutation instance. func NewPermutation(t, rf, rp int) *Permutation { if t < 2 || t > 3 { diff --git a/ecc/bw6-633/fr/poseidon2/poseidon2_test.go b/ecc/bw6-633/fr/poseidon2/poseidon2_test.go index 06de110947..4bb1318954 100644 --- a/ecc/bw6-633/fr/poseidon2/poseidon2_test.go +++ b/ecc/bw6-633/fr/poseidon2/poseidon2_test.go @@ -109,3 +109,49 @@ func TestHashReset(t *testing.T) { require.Equal(t, res, h.Sum(nil)) } + +// denseMatrixMul returns m·x. +func denseMatrixMul(m [][]fr.Element, x []fr.Element) []fr.Element { + res := make([]fr.Element, len(m)) + for i := range m { + var tmp fr.Element + for j := range m[i] { + tmp.Mul(&m[i][j], &x[j]) + res[i].Add(&res[i], &tmp) + } + } + return res +} + +// TestDenseMatricesMatchInPlace checks that the dense matrices act on a random +// state as the in-place multiplications of the permutation. +func TestDenseMatricesMatchInPlace(t *testing.T) { + for _, tc := range []struct { + width, nbFullRounds, nbPartialRounds int + }{ + {2, 8, 56}, + {3, 8, 56}, + } { + h := NewPermutation(tc.width, tc.nbFullRounds, tc.nbPartialRounds) + + for name, step := range map[string]struct { + dense [][]fr.Element + inPlace func([]fr.Element) + }{ + "external": {h.params.ExternalMatrix(), h.matMulExternalInPlace}, + "internal": {h.params.InternalMatrix(), h.matMulInternalInPlace}, + } { + x := make([]fr.Element, tc.width) + for i := range x { + x[i].MustSetRandom() + } + want := denseMatrixMul(step.dense, x) + step.inPlace(x) + for i := range x { + if !x[i].Equal(&want[i]) { + t.Fatalf("width %d, %s matrix: dense and in-place multiplications differ at index %d", tc.width, name, i) + } + } + } + } +} diff --git a/ecc/bw6-761/fr/poseidon2/poseidon2.go b/ecc/bw6-761/fr/poseidon2/poseidon2.go index ee908ff421..8ae5a83a66 100644 --- a/ecc/bw6-761/fr/poseidon2/poseidon2.go +++ b/ecc/bw6-761/fr/poseidon2/poseidon2.go @@ -12,6 +12,7 @@ import ( "golang.org/x/crypto/sha3" "github.com/consensys/gnark-crypto/ecc/bw6-761/fr" + internalposeidon2 "github.com/consensys/gnark-crypto/internal/poseidon2" ) var ( @@ -121,6 +122,43 @@ type Permutation struct { params *Parameters } +// ExternalMatrix returns the dense external matrix M_E of the permutation, which +// is applied to the state before the first round and after each full round. +// With I and J the identity and the all-ones matrices, and M4 the 4×4 block of +// the Poseidon2 paper (rows 5 7 1 3 / 4 6 1 1 / 1 3 5 7 / 1 1 4 6), M_E is +// +// - I + J for width 2 and 3; +// - M4 for width 4; +// - circ(2·M4, M4, ..., M4) = (I + J) ⊗ M4 for width 4k with k > 1. +func (p *Parameters) ExternalMatrix() [][]fr.Element { + return internalposeidon2.ExternalMatrix[fr.Element](p.Width, internalposeidon2.Paper) +} + +// InternalMatrix returns the dense internal matrix M_I of the permutation, which +// is applied to the state after each partial round. With J the all-ones matrix, +// +// M_I = J + diag(d), +// +// that is, for a state x, (M_I·x)_i = Σ_j x_j + d_i·x_i, where d is the diagonal +// of the permutation. +func (p *Parameters) InternalMatrix() [][]fr.Element { + var d []fr.Element + switch p.Width { + case 2: + d = make([]fr.Element, 2) + d[0].SetOne() + d[1].SetUint64(2) + case 3: + d = make([]fr.Element, 3) + d[0].SetOne() + d[1].SetOne() + d[2].SetUint64(2) + default: + panic("only Width=2,3 are supported") + } + return internalposeidon2.InternalMatrix(d) +} + // NewPermutation returns a new Poseidon2 permutation instance. func NewPermutation(t, rf, rp int) *Permutation { if t < 2 || t > 3 { diff --git a/ecc/bw6-761/fr/poseidon2/poseidon2_test.go b/ecc/bw6-761/fr/poseidon2/poseidon2_test.go index 9e3a52cd7c..5264feb873 100644 --- a/ecc/bw6-761/fr/poseidon2/poseidon2_test.go +++ b/ecc/bw6-761/fr/poseidon2/poseidon2_test.go @@ -109,3 +109,49 @@ func TestHashReset(t *testing.T) { require.Equal(t, res, h.Sum(nil)) } + +// denseMatrixMul returns m·x. +func denseMatrixMul(m [][]fr.Element, x []fr.Element) []fr.Element { + res := make([]fr.Element, len(m)) + for i := range m { + var tmp fr.Element + for j := range m[i] { + tmp.Mul(&m[i][j], &x[j]) + res[i].Add(&res[i], &tmp) + } + } + return res +} + +// TestDenseMatricesMatchInPlace checks that the dense matrices act on a random +// state as the in-place multiplications of the permutation. +func TestDenseMatricesMatchInPlace(t *testing.T) { + for _, tc := range []struct { + width, nbFullRounds, nbPartialRounds int + }{ + {2, 8, 56}, + {3, 8, 56}, + } { + h := NewPermutation(tc.width, tc.nbFullRounds, tc.nbPartialRounds) + + for name, step := range map[string]struct { + dense [][]fr.Element + inPlace func([]fr.Element) + }{ + "external": {h.params.ExternalMatrix(), h.matMulExternalInPlace}, + "internal": {h.params.InternalMatrix(), h.matMulInternalInPlace}, + } { + x := make([]fr.Element, tc.width) + for i := range x { + x[i].MustSetRandom() + } + want := denseMatrixMul(step.dense, x) + step.inPlace(x) + for i := range x { + if !x[i].Equal(&want[i]) { + t.Fatalf("width %d, %s matrix: dense and in-place multiplications differ at index %d", tc.width, name, i) + } + } + } + } +} diff --git a/ecc/grumpkin/fr/poseidon2/poseidon2.go b/ecc/grumpkin/fr/poseidon2/poseidon2.go index 2aaeef27fa..aba28edc08 100644 --- a/ecc/grumpkin/fr/poseidon2/poseidon2.go +++ b/ecc/grumpkin/fr/poseidon2/poseidon2.go @@ -12,6 +12,7 @@ import ( "golang.org/x/crypto/sha3" "github.com/consensys/gnark-crypto/ecc/grumpkin/fr" + internalposeidon2 "github.com/consensys/gnark-crypto/internal/poseidon2" ) var ( @@ -121,6 +122,43 @@ type Permutation struct { params *Parameters } +// ExternalMatrix returns the dense external matrix M_E of the permutation, which +// is applied to the state before the first round and after each full round. +// With I and J the identity and the all-ones matrices, and M4 the 4×4 block of +// the Poseidon2 paper (rows 5 7 1 3 / 4 6 1 1 / 1 3 5 7 / 1 1 4 6), M_E is +// +// - I + J for width 2 and 3; +// - M4 for width 4; +// - circ(2·M4, M4, ..., M4) = (I + J) ⊗ M4 for width 4k with k > 1. +func (p *Parameters) ExternalMatrix() [][]fr.Element { + return internalposeidon2.ExternalMatrix[fr.Element](p.Width, internalposeidon2.Paper) +} + +// InternalMatrix returns the dense internal matrix M_I of the permutation, which +// is applied to the state after each partial round. With J the all-ones matrix, +// +// M_I = J + diag(d), +// +// that is, for a state x, (M_I·x)_i = Σ_j x_j + d_i·x_i, where d is the diagonal +// of the permutation. +func (p *Parameters) InternalMatrix() [][]fr.Element { + var d []fr.Element + switch p.Width { + case 2: + d = make([]fr.Element, 2) + d[0].SetOne() + d[1].SetUint64(2) + case 3: + d = make([]fr.Element, 3) + d[0].SetOne() + d[1].SetOne() + d[2].SetUint64(2) + default: + panic("only Width=2,3 are supported") + } + return internalposeidon2.InternalMatrix(d) +} + // NewPermutation returns a new Poseidon2 permutation instance. func NewPermutation(t, rf, rp int) *Permutation { if t < 2 || t > 3 { diff --git a/ecc/grumpkin/fr/poseidon2/poseidon2_test.go b/ecc/grumpkin/fr/poseidon2/poseidon2_test.go index d9e6b89c28..653198594d 100644 --- a/ecc/grumpkin/fr/poseidon2/poseidon2_test.go +++ b/ecc/grumpkin/fr/poseidon2/poseidon2_test.go @@ -109,3 +109,49 @@ func TestHashReset(t *testing.T) { require.Equal(t, res, h.Sum(nil)) } + +// denseMatrixMul returns m·x. +func denseMatrixMul(m [][]fr.Element, x []fr.Element) []fr.Element { + res := make([]fr.Element, len(m)) + for i := range m { + var tmp fr.Element + for j := range m[i] { + tmp.Mul(&m[i][j], &x[j]) + res[i].Add(&res[i], &tmp) + } + } + return res +} + +// TestDenseMatricesMatchInPlace checks that the dense matrices act on a random +// state as the in-place multiplications of the permutation. +func TestDenseMatricesMatchInPlace(t *testing.T) { + for _, tc := range []struct { + width, nbFullRounds, nbPartialRounds int + }{ + {2, 8, 56}, + {3, 8, 56}, + } { + h := NewPermutation(tc.width, tc.nbFullRounds, tc.nbPartialRounds) + + for name, step := range map[string]struct { + dense [][]fr.Element + inPlace func([]fr.Element) + }{ + "external": {h.params.ExternalMatrix(), h.matMulExternalInPlace}, + "internal": {h.params.InternalMatrix(), h.matMulInternalInPlace}, + } { + x := make([]fr.Element, tc.width) + for i := range x { + x[i].MustSetRandom() + } + want := denseMatrixMul(step.dense, x) + step.inPlace(x) + for i := range x { + if !x[i].Equal(&want[i]) { + t.Fatalf("width %d, %s matrix: dense and in-place multiplications differ at index %d", tc.width, name, i) + } + } + } + } +} diff --git a/field/babybear/poseidon2/poseidon2.go b/field/babybear/poseidon2/poseidon2.go index d303d6e414..2a3b1a18f8 100644 --- a/field/babybear/poseidon2/poseidon2.go +++ b/field/babybear/poseidon2/poseidon2.go @@ -16,6 +16,7 @@ import ( "golang.org/x/crypto/sha3" fr "github.com/consensys/gnark-crypto/field/babybear" + internalposeidon2 "github.com/consensys/gnark-crypto/internal/poseidon2" "github.com/consensys/gnark-crypto/utils/cpu" ) @@ -130,6 +131,36 @@ type Permutation struct { params *Parameters } +// ExternalMatrix returns the dense external matrix M_E of the permutation, which +// is applied to the state before the first round and after each full round. +// With I and J the identity and the all-ones matrices, and M4 the 4×4 block of +// Plonky3 (rows 2 3 1 1 / 1 2 3 1 / 1 1 2 3 / 3 1 1 2), M_E is +// +// - I + J for width 2 and 3; +// - M4 for width 4; +// - circ(2·M4, M4, ..., M4) = (I + J) ⊗ M4 for width 4k with k > 1. +func (p *Parameters) ExternalMatrix() [][]fr.Element { + return internalposeidon2.ExternalMatrix[fr.Element](p.Width, internalposeidon2.Plonky3) +} + +// InternalMatrix returns the dense internal matrix M_I of the permutation, which +// is applied to the state after each partial round. With J the all-ones matrix, +// +// M_I = J + diag(d), +// +// that is, for a state x, (M_I·x)_i = Σ_j x_j + d_i·x_i, where d is the diagonal +// of the permutation (diag16 for width 16, diag24 for width 24). +func (p *Parameters) InternalMatrix() [][]fr.Element { + switch p.Width { + case 16: + return internalposeidon2.InternalMatrix(diag16) + case 24: + return internalposeidon2.InternalMatrix(diag24) + default: + panic("only Width=16,24 are supported") + } +} + // NewPermutation returns a new Poseidon2 permutation instance. func NewPermutation(t, rf, rp int) *Permutation { if t != 16 && t != 24 { diff --git a/field/babybear/poseidon2/poseidon2_test.go b/field/babybear/poseidon2/poseidon2_test.go index e2037cb4a3..5317b63e7e 100644 --- a/field/babybear/poseidon2/poseidon2_test.go +++ b/field/babybear/poseidon2/poseidon2_test.go @@ -63,6 +63,52 @@ func TestMulMulInternalInPlaceWidth24(t *testing.T) { } } +// denseMatrixMul returns m·x. +func denseMatrixMul(m [][]fr.Element, x []fr.Element) []fr.Element { + res := make([]fr.Element, len(m)) + for i := range m { + var tmp fr.Element + for j := range m[i] { + tmp.Mul(&m[i][j], &x[j]) + res[i].Add(&res[i], &tmp) + } + } + return res +} + +// TestDenseMatricesMatchInPlace checks that the dense matrices act on a random +// state as the in-place multiplications of the permutation. +func TestDenseMatricesMatchInPlace(t *testing.T) { + for _, tc := range []struct { + width, nbFullRounds, nbPartialRounds int + }{ + {16, 8, 13}, + {24, 8, 21}, + } { + h := NewPermutation(tc.width, tc.nbFullRounds, tc.nbPartialRounds) + + for name, step := range map[string]struct { + dense [][]fr.Element + inPlace func([]fr.Element) + }{ + "external": {h.params.ExternalMatrix(), h.matMulExternalInPlace}, + "internal": {h.params.InternalMatrix(), h.matMulInternalInPlace}, + } { + x := make([]fr.Element, tc.width) + for i := range x { + x[i].MustSetRandom() + } + want := denseMatrixMul(step.dense, x) + step.inPlace(x) + for i := range x { + if !x[i].Equal(&want[i]) { + t.Fatalf("width %d, %s matrix: dense and in-place multiplications differ at index %d", tc.width, name, i) + } + } + } + } +} + func TestAVX512Width16(t *testing.T) { if !cpu.SupportAVX512 { t.Skip("AVX512 not supported") diff --git a/field/goldilocks/poseidon2/poseidon2.go b/field/goldilocks/poseidon2/poseidon2.go index dbd15b2a56..b76334077d 100644 --- a/field/goldilocks/poseidon2/poseidon2.go +++ b/field/goldilocks/poseidon2/poseidon2.go @@ -15,6 +15,7 @@ import ( "golang.org/x/crypto/sha3" fr "github.com/consensys/gnark-crypto/field/goldilocks" + internalposeidon2 "github.com/consensys/gnark-crypto/internal/poseidon2" ) var ( @@ -121,6 +122,36 @@ type Permutation struct { params *Parameters } +// ExternalMatrix returns the dense external matrix M_E of the permutation, which +// is applied to the state before the first round and after each full round. +// With I and J the identity and the all-ones matrices, and M4 the 4×4 block of +// the Poseidon2 paper (rows 5 7 1 3 / 4 6 1 1 / 1 3 5 7 / 1 1 4 6), M_E is +// +// - I + J for width 2 and 3; +// - M4 for width 4; +// - circ(2·M4, M4, ..., M4) = (I + J) ⊗ M4 for width 4k with k > 1. +func (p *Parameters) ExternalMatrix() [][]fr.Element { + return internalposeidon2.ExternalMatrix[fr.Element](p.Width, internalposeidon2.Paper) +} + +// InternalMatrix returns the dense internal matrix M_I of the permutation, which +// is applied to the state after each partial round. With J the all-ones matrix, +// +// M_I = J + diag(d), +// +// that is, for a state x, (M_I·x)_i = Σ_j x_j + d_i·x_i, where d is the diagonal +// of the permutation (diag8 for width 8, diag12 for width 12). +func (p *Parameters) InternalMatrix() [][]fr.Element { + switch p.Width { + case 8: + return internalposeidon2.InternalMatrix(diag8) + case 12: + return internalposeidon2.InternalMatrix(diag12) + default: + panic("only Width=8,12 are supported") + } +} + // NewPermutation returns a new Poseidon2 permutation instance. func NewPermutation(t, rf, rp int) *Permutation { if t != 8 && t != 12 { diff --git a/field/goldilocks/poseidon2/poseidon2_test.go b/field/goldilocks/poseidon2/poseidon2_test.go index 1dfe183054..a4be7c58a0 100644 --- a/field/goldilocks/poseidon2/poseidon2_test.go +++ b/field/goldilocks/poseidon2/poseidon2_test.go @@ -60,6 +60,52 @@ func TestMulMulInternalInPlaceWidth12(t *testing.T) { } } } + +// denseMatrixMul returns m·x. +func denseMatrixMul(m [][]fr.Element, x []fr.Element) []fr.Element { + res := make([]fr.Element, len(m)) + for i := range m { + var tmp fr.Element + for j := range m[i] { + tmp.Mul(&m[i][j], &x[j]) + res[i].Add(&res[i], &tmp) + } + } + return res +} + +// TestDenseMatricesMatchInPlace checks that the dense matrices act on a random +// state as the in-place multiplications of the permutation. +func TestDenseMatricesMatchInPlace(t *testing.T) { + for _, tc := range []struct { + width, nbFullRounds, nbPartialRounds int + }{ + {8, 6, 17}, + {12, 6, 17}, + } { + h := NewPermutation(tc.width, tc.nbFullRounds, tc.nbPartialRounds) + + for name, step := range map[string]struct { + dense [][]fr.Element + inPlace func([]fr.Element) + }{ + "external": {h.params.ExternalMatrix(), h.matMulExternalInPlace}, + "internal": {h.params.InternalMatrix(), h.matMulInternalInPlace}, + } { + x := make([]fr.Element, tc.width) + for i := range x { + x[i].MustSetRandom() + } + want := denseMatrixMul(step.dense, x) + step.inPlace(x) + for i := range x { + if !x[i].Equal(&want[i]) { + t.Fatalf("width %d, %s matrix: dense and in-place multiplications differ at index %d", tc.width, name, i) + } + } + } + } +} func TestPoseidon2Width8(t *testing.T) { var input, expected [8]fr.Element // these are random values generated by MustSetRandom() diff --git a/field/koalabear/poseidon2/poseidon2.go b/field/koalabear/poseidon2/poseidon2.go index abaa52da5c..c6f79cd547 100644 --- a/field/koalabear/poseidon2/poseidon2.go +++ b/field/koalabear/poseidon2/poseidon2.go @@ -16,6 +16,7 @@ import ( "golang.org/x/crypto/sha3" fr "github.com/consensys/gnark-crypto/field/koalabear" + internalposeidon2 "github.com/consensys/gnark-crypto/internal/poseidon2" "github.com/consensys/gnark-crypto/utils/cpu" ) @@ -130,6 +131,36 @@ type Permutation struct { params *Parameters } +// ExternalMatrix returns the dense external matrix M_E of the permutation, which +// is applied to the state before the first round and after each full round. +// With I and J the identity and the all-ones matrices, and M4 the 4×4 block of +// Plonky3 (rows 2 3 1 1 / 1 2 3 1 / 1 1 2 3 / 3 1 1 2), M_E is +// +// - I + J for width 2 and 3; +// - M4 for width 4; +// - circ(2·M4, M4, ..., M4) = (I + J) ⊗ M4 for width 4k with k > 1. +func (p *Parameters) ExternalMatrix() [][]fr.Element { + return internalposeidon2.ExternalMatrix[fr.Element](p.Width, internalposeidon2.Plonky3) +} + +// InternalMatrix returns the dense internal matrix M_I of the permutation, which +// is applied to the state after each partial round. With J the all-ones matrix, +// +// M_I = J + diag(d), +// +// that is, for a state x, (M_I·x)_i = Σ_j x_j + d_i·x_i, where d is the diagonal +// of the permutation (diag16 for width 16, diag24 for width 24). +func (p *Parameters) InternalMatrix() [][]fr.Element { + switch p.Width { + case 16: + return internalposeidon2.InternalMatrix(diag16) + case 24: + return internalposeidon2.InternalMatrix(diag24) + default: + panic("only Width=16,24 are supported") + } +} + // NewPermutation returns a new Poseidon2 permutation instance. func NewPermutation(t, rf, rp int) *Permutation { if t != 16 && t != 24 { diff --git a/field/koalabear/poseidon2/poseidon2_test.go b/field/koalabear/poseidon2/poseidon2_test.go index c47680863f..f0363fe834 100644 --- a/field/koalabear/poseidon2/poseidon2_test.go +++ b/field/koalabear/poseidon2/poseidon2_test.go @@ -64,6 +64,52 @@ func TestMulMulInternalInPlaceWidth24(t *testing.T) { } } +// denseMatrixMul returns m·x. +func denseMatrixMul(m [][]fr.Element, x []fr.Element) []fr.Element { + res := make([]fr.Element, len(m)) + for i := range m { + var tmp fr.Element + for j := range m[i] { + tmp.Mul(&m[i][j], &x[j]) + res[i].Add(&res[i], &tmp) + } + } + return res +} + +// TestDenseMatricesMatchInPlace checks that the dense matrices act on a random +// state as the in-place multiplications of the permutation. +func TestDenseMatricesMatchInPlace(t *testing.T) { + for _, tc := range []struct { + width, nbFullRounds, nbPartialRounds int + }{ + {16, 6, 21}, + {24, 6, 21}, + } { + h := NewPermutation(tc.width, tc.nbFullRounds, tc.nbPartialRounds) + + for name, step := range map[string]struct { + dense [][]fr.Element + inPlace func([]fr.Element) + }{ + "external": {h.params.ExternalMatrix(), h.matMulExternalInPlace}, + "internal": {h.params.InternalMatrix(), h.matMulInternalInPlace}, + } { + x := make([]fr.Element, tc.width) + for i := range x { + x[i].MustSetRandom() + } + want := denseMatrixMul(step.dense, x) + step.inPlace(x) + for i := range x { + if !x[i].Equal(&want[i]) { + t.Fatalf("width %d, %s matrix: dense and in-place multiplications differ at index %d", tc.width, name, i) + } + } + } + } +} + func TestAVX512Width16(t *testing.T) { if !cpu.SupportAVX512 { t.Skip("AVX512 not supported") diff --git a/field/mamabear/poseidon2/poseidon2.go b/field/mamabear/poseidon2/poseidon2.go index 278bd39f60..ecb25f7601 100644 --- a/field/mamabear/poseidon2/poseidon2.go +++ b/field/mamabear/poseidon2/poseidon2.go @@ -15,6 +15,7 @@ import ( "golang.org/x/crypto/sha3" fr "github.com/consensys/gnark-crypto/field/mamabear" + internalposeidon2 "github.com/consensys/gnark-crypto/internal/poseidon2" ) var ( @@ -121,6 +122,36 @@ type Permutation struct { params *Parameters } +// ExternalMatrix returns the dense external matrix M_E of the permutation, which +// is applied to the state before the first round and after each full round. +// With I and J the identity and the all-ones matrices, and M4 the 4×4 block of +// Plonky3 (rows 2 3 1 1 / 1 2 3 1 / 1 1 2 3 / 3 1 1 2), M_E is +// +// - I + J for width 2 and 3; +// - M4 for width 4; +// - circ(2·M4, M4, ..., M4) = (I + J) ⊗ M4 for width 4k with k > 1. +func (p *Parameters) ExternalMatrix() [][]fr.Element { + return internalposeidon2.ExternalMatrix[fr.Element](p.Width, internalposeidon2.Plonky3) +} + +// InternalMatrix returns the dense internal matrix M_I of the permutation, which +// is applied to the state after each partial round. With J the all-ones matrix, +// +// M_I = J + diag(d), +// +// that is, for a state x, (M_I·x)_i = Σ_j x_j + d_i·x_i, where d is the diagonal +// of the permutation (diag16 for width 16, diag24 for width 24). +func (p *Parameters) InternalMatrix() [][]fr.Element { + switch p.Width { + case 16: + return internalposeidon2.InternalMatrix(diag16) + case 24: + return internalposeidon2.InternalMatrix(diag24) + default: + panic("only Width=16,24 are supported") + } +} + // NewPermutation returns a new Poseidon2 permutation instance. func NewPermutation(t, rf, rp int) *Permutation { if t != 16 && t != 24 { diff --git a/field/mamabear/poseidon2/poseidon2_test.go b/field/mamabear/poseidon2/poseidon2_test.go index de9cb20624..a7da662715 100644 --- a/field/mamabear/poseidon2/poseidon2_test.go +++ b/field/mamabear/poseidon2/poseidon2_test.go @@ -61,6 +61,52 @@ func TestMulMulInternalInPlaceWidth24(t *testing.T) { } } +// denseMatrixMul returns m·x. +func denseMatrixMul(m [][]fr.Element, x []fr.Element) []fr.Element { + res := make([]fr.Element, len(m)) + for i := range m { + var tmp fr.Element + for j := range m[i] { + tmp.Mul(&m[i][j], &x[j]) + res[i].Add(&res[i], &tmp) + } + } + return res +} + +// TestDenseMatricesMatchInPlace checks that the dense matrices act on a random +// state as the in-place multiplications of the permutation. +func TestDenseMatricesMatchInPlace(t *testing.T) { + for _, tc := range []struct { + width, nbFullRounds, nbPartialRounds int + }{ + {16, 6, 21}, + {24, 6, 21}, + } { + h := NewPermutation(tc.width, tc.nbFullRounds, tc.nbPartialRounds) + + for name, step := range map[string]struct { + dense [][]fr.Element + inPlace func([]fr.Element) + }{ + "external": {h.params.ExternalMatrix(), h.matMulExternalInPlace}, + "internal": {h.params.InternalMatrix(), h.matMulInternalInPlace}, + } { + x := make([]fr.Element, tc.width) + for i := range x { + x[i].MustSetRandom() + } + want := denseMatrixMul(step.dense, x) + step.inPlace(x) + for i := range x { + if !x[i].Equal(&want[i]) { + t.Fatalf("width %d, %s matrix: dense and in-place multiplications differ at index %d", tc.width, name, i) + } + } + } + } +} + func TestPoseidon2Width16(t *testing.T) { var input, expected [16]fr.Element // these are random values generated by MustSetRandom() diff --git a/internal/generator/crypto/hash/poseidon2/template/poseidon2.go.tmpl b/internal/generator/crypto/hash/poseidon2/template/poseidon2.go.tmpl index e7c0ff37f3..0bab8fd649 100644 --- a/internal/generator/crypto/hash/poseidon2/template/poseidon2.go.tmpl +++ b/internal/generator/crypto/hash/poseidon2/template/poseidon2.go.tmpl @@ -5,6 +5,7 @@ import ( "golang.org/x/crypto/sha3" "github.com/consensys/gnark-crypto/ecc/{{ .Name }}/fr" + internalposeidon2 "github.com/consensys/gnark-crypto/internal/poseidon2" ) var ( @@ -147,6 +148,50 @@ type Permutation struct { params *Parameters } +// ExternalMatrix returns the dense external matrix M_E of the permutation, which +// is applied to the state before the first round and after each full round. +// With I and J the identity and the all-ones matrices, and M4 the 4×4 block of +// the Poseidon2 paper (rows 5 7 1 3 / 4 6 1 1 / 1 3 5 7 / 1 1 4 6), M_E is +// +// - I + J for width 2 and 3; +// - M4 for width 4; +// - circ(2·M4, M4, ..., M4) = (I + J) ⊗ M4 for width 4k with k > 1. +func (p *Parameters) ExternalMatrix() [][]fr.Element { + return internalposeidon2.ExternalMatrix[fr.Element](p.Width, internalposeidon2.Paper) +} + +// InternalMatrix returns the dense internal matrix M_I of the permutation, which +// is applied to the state after each partial round. With J the all-ones matrix, +// +// M_I = J + diag(d), +// +// that is, for a state x, (M_I·x)_i = Σ_j x_j + d_i·x_i, where d is the diagonal +// of the permutation{{- if eq .Name "bn254" }} (DiagM1 for the widths of 4 and more){{ end -}}. +func (p *Parameters) InternalMatrix() [][]fr.Element { + var d []fr.Element + switch p.Width { + case 2: + d = make([]fr.Element, 2) + d[0].SetOne() + d[1].SetUint64(2) + case 3: + d = make([]fr.Element, 3) + d[0].SetOne() + d[1].SetOne() + d[2].SetUint64(2) + default: + {{- if eq .Name "bn254" }} + if len(p.DiagM1) != p.Width { + panic("missing internal matrix diagonal") + } + d = p.DiagM1 + {{- else }} + panic("only Width=2,3 are supported") + {{- end }} + } + return internalposeidon2.InternalMatrix(d) +} + // NewPermutation returns a new Poseidon2 permutation instance. func NewPermutation(t, rf, rp int) *Permutation { {{- if eq .Name "bn254" }} diff --git a/internal/generator/crypto/hash/poseidon2/template/poseidon2.test.go.tmpl b/internal/generator/crypto/hash/poseidon2/template/poseidon2.test.go.tmpl index 115c22de9d..46675afd29 100644 --- a/internal/generator/crypto/hash/poseidon2/template/poseidon2.test.go.tmpl +++ b/internal/generator/crypto/hash/poseidon2/template/poseidon2.test.go.tmpl @@ -232,4 +232,56 @@ func TestHashReset(t *testing.T) { require.NoError(t, err) require.Equal(t, res, h.Sum(nil)) -} \ No newline at end of file +} + +// denseMatrixMul returns m·x. +func denseMatrixMul(m [][]fr.Element, x []fr.Element) []fr.Element { + res := make([]fr.Element, len(m)) + for i := range m { + var tmp fr.Element + for j := range m[i] { + tmp.Mul(&m[i][j], &x[j]) + res[i].Add(&res[i], &tmp) + } + } + return res +} + +// TestDenseMatricesMatchInPlace checks that the dense matrices act on a random +// state as the in-place multiplications of the permutation. +func TestDenseMatricesMatchInPlace(t *testing.T) { + for _, tc := range []struct { + width, nbFullRounds, nbPartialRounds int + }{ + {2, 8, 56}, + {3, 8, 56}, +{{- if eq .Name "bn254" }} + {4, 8, 56}, + {8, 8, 57}, + {12, 8, 57}, + {16, 8, 57}, +{{- end }} + } { + h := NewPermutation(tc.width, tc.nbFullRounds, tc.nbPartialRounds) + + for name, step := range map[string]struct { + dense [][]fr.Element + inPlace func([]fr.Element) + }{ + "external": {h.params.ExternalMatrix(), h.matMulExternalInPlace}, + "internal": {h.params.InternalMatrix(), h.matMulInternalInPlace}, + } { + x := make([]fr.Element, tc.width) + for i := range x { + x[i].MustSetRandom() + } + want := denseMatrixMul(step.dense, x) + step.inPlace(x) + for i := range x { + if !x[i].Equal(&want[i]) { + t.Fatalf("width %d, %s matrix: dense and in-place multiplications differ at index %d", tc.width, name, i) + } + } + } + } +} diff --git a/internal/generator/field/template/poseidon2/poseidon2.go.tmpl b/internal/generator/field/template/poseidon2/poseidon2.go.tmpl index 7c3c2aef1c..f41eaacee2 100644 --- a/internal/generator/field/template/poseidon2/poseidon2.go.tmpl +++ b/internal/generator/field/template/poseidon2/poseidon2.go.tmpl @@ -9,6 +9,7 @@ import ( "golang.org/x/crypto/sha3" fr "{{ .FieldPackagePath }}" + internalposeidon2 "github.com/consensys/gnark-crypto/internal/poseidon2" {{- if .F31}} "github.com/consensys/gnark-crypto/utils/cpu" @@ -145,6 +146,36 @@ type Permutation struct { params *Parameters } +// ExternalMatrix returns the dense external matrix M_E of the permutation, which +// is applied to the state before the first round and after each full round. +// With I and J the identity and the all-ones matrices, and M4 the 4×4 block of +// {{if eq .FF "goldilocks"}}the Poseidon2 paper (rows 5 7 1 3 / 4 6 1 1 / 1 3 5 7 / 1 1 4 6){{else}}Plonky3 (rows 2 3 1 1 / 1 2 3 1 / 1 1 2 3 / 3 1 1 2){{end}}, M_E is +// +// - I + J for width 2 and 3; +// - M4 for width 4; +// - circ(2·M4, M4, ..., M4) = (I + J) ⊗ M4 for width 4k with k > 1. +func (p *Parameters) ExternalMatrix() [][]fr.Element { + return internalposeidon2.ExternalMatrix[fr.Element](p.Width, internalposeidon2.{{if eq .FF "goldilocks"}}Paper{{else}}Plonky3{{end}}) +} + +// InternalMatrix returns the dense internal matrix M_I of the permutation, which +// is applied to the state after each partial round. With J the all-ones matrix, +// +// M_I = J + diag(d), +// +// that is, for a state x, (M_I·x)_i = Σ_j x_j + d_i·x_i, where d is the diagonal +// of the permutation (diag{{$wc}} for width {{$wc}}, diag{{$ws}} for width {{$ws}}). +func (p *Parameters) InternalMatrix() [][]fr.Element { + switch p.Width { + case {{$wc}}: + return internalposeidon2.InternalMatrix(diag{{$wc}}) + case {{$ws}}: + return internalposeidon2.InternalMatrix(diag{{$ws}}) + default: + panic("only Width={{$wc}},{{$ws}} are supported") + } +} + // NewPermutation returns a new Poseidon2 permutation instance. func NewPermutation(t, rf, rp int) *Permutation { {{- if or (eq .Family "babybear") (eq .Family "koalabear")}} diff --git a/internal/generator/field/template/poseidon2/poseidon2_test.go.tmpl b/internal/generator/field/template/poseidon2/poseidon2_test.go.tmpl index 56ac0af967..1e031b04db 100644 --- a/internal/generator/field/template/poseidon2/poseidon2_test.go.tmpl +++ b/internal/generator/field/template/poseidon2/poseidon2_test.go.tmpl @@ -68,6 +68,53 @@ func TestMulMulInternalInPlaceWidth{{- $w1}}(t *testing.T) { } +// denseMatrixMul returns m·x. +func denseMatrixMul(m [][]fr.Element, x []fr.Element) []fr.Element { + res := make([]fr.Element, len(m)) + for i := range m { + var tmp fr.Element + for j := range m[i] { + tmp.Mul(&m[i][j], &x[j]) + res[i].Add(&res[i], &tmp) + } + } + return res +} + +// TestDenseMatricesMatchInPlace checks that the dense matrices act on a random +// state as the in-place multiplications of the permutation. +func TestDenseMatricesMatchInPlace(t *testing.T) { + for _, tc := range []struct { + width, nbFullRounds, nbPartialRounds int + }{ + { {{- $w0}}, {{.ParamsCompression.FullRounds}}, {{.ParamsCompression.PartialRounds}} }, + { {{- $w1}}, {{.ParamsSponge.FullRounds}}, {{.ParamsSponge.PartialRounds}} }, + } { + h := NewPermutation(tc.width, tc.nbFullRounds, tc.nbPartialRounds) + + for name, step := range map[string]struct { + dense [][]fr.Element + inPlace func([]fr.Element) + }{ + "external": {h.params.ExternalMatrix(), h.matMulExternalInPlace}, + "internal": {h.params.InternalMatrix(), h.matMulInternalInPlace}, + } { + x := make([]fr.Element, tc.width) + for i := range x { + x[i].MustSetRandom() + } + want := denseMatrixMul(step.dense, x) + step.inPlace(x) + for i := range x { + if !x[i].Equal(&want[i]) { + t.Fatalf("width %d, %s matrix: dense and in-place multiplications differ at index %d", tc.width, name, i) + } + } + } + } +} + + {{- if .F31}} diff --git a/internal/poseidon2/matrices.go b/internal/poseidon2/matrices.go new file mode 100644 index 0000000000..2129cf7fcc --- /dev/null +++ b/internal/poseidon2/matrices.go @@ -0,0 +1,121 @@ +// Copyright 2020-2026 Consensys Software Inc. +// Licensed under the Apache License, Version 2.0. See the LICENSE file for details. + +// Package poseidon2 holds the definitions of the linear layers of the Poseidon2 +// permutation that are shared by the generated field and curve packages. +package poseidon2 + +// ExternalMatrixKind names the source of the 4×4 block M4 from which the +// external matrix is built for the widths that are multiples of 4. +type ExternalMatrixKind int + +const ( + // Plonky3 is the block used by Plonky3, whose rows are + // + // 2 3 1 1 + // 1 2 3 1 + // 1 1 2 3 + // 3 1 1 2 + Plonky3 ExternalMatrixKind = iota + // Paper is the block of the Poseidon2 paper (https://eprint.iacr.org/2023/323.pdf, + // appendix B), whose rows are + // + // 5 7 1 3 + // 4 6 1 1 + // 1 3 5 7 + // 1 1 4 6 + Paper +) + +// m4Blocks holds the block M4 of each ExternalMatrixKind. +var m4Blocks = [...][4][4]int64{ + Plonky3: { + {2, 3, 1, 1}, + {1, 2, 3, 1}, + {1, 1, 2, 3}, + {3, 1, 1, 2}, + }, + Paper: { + {5, 7, 1, 3}, + {4, 6, 1, 1}, + {1, 3, 5, 7}, + {1, 1, 4, 6}, + }, +} + +// ring is the set of operations on a pointer PE to a ring element E that the +// matrix constructions need. +type ring[E any] interface { + *E + SetInt64(int64) *E + Add(*E, *E) *E +} + +// ExternalMatrix returns the dense width×width external matrix M_E, which the +// permutation applies before the first round and after each full round. With I +// and J the identity and the all-ones matrices, it is +// +// - I + J for width 2 and 3, whatever the kind; +// - M4 for width 4; +// - circ(2·M4, M4, ..., M4) = (I + J) ⊗ M4 for width 4k with k > 1, where +// I and J have size k. +// +// M4 is the block selected by kind. It panics for any other width, and for an +// unknown kind when a block is needed. +func ExternalMatrix[E any, PE ring[E]](width int, kind ExternalMatrixKind) [][]E { + m := newMatrix[E](width) + + switch { + case width == 2 || width == 3: + for i := range m { + for j := range m[i] { + v := int64(1) + if i == j { + v = 2 + } + PE(&m[i][j]).SetInt64(v) + } + } + return m + case width%4 != 0 || width <= 0: + panic("poseidon2: only widths 2, 3 and multiples of 4 are supported") + } + + if kind < 0 || int(kind) >= len(m4Blocks) { + panic("poseidon2: unknown external matrix kind") + } + m4 := &m4Blocks[kind] + for i := range m { + for j := range m[i] { + v := m4[i%4][j%4] + if width > 4 && i/4 == j/4 { + v *= 2 + } + PE(&m[i][j]).SetInt64(v) + } + } + return m +} + +// InternalMatrix returns the dense internal matrix M_I = J + diag(d) of width +// len(d), where J is the all-ones matrix: entry (i, j) is 1, plus d[i] when +// i == j. For a state x, (M_I·x)_i = Σ_j x_j + d[i]·x_i. The permutation applies +// M_I after each partial round. +func InternalMatrix[E any, PE ring[E]](d []E) [][]E { + m := newMatrix[E](len(d)) + for i := range m { + for j := range m[i] { + PE(&m[i][j]).SetInt64(1) + } + PE(&m[i][i]).Add(PE(&m[i][i]), PE(&d[i])) + } + return m +} + +func newMatrix[E any](n int) [][]E { + m := make([][]E, n) + for i := range m { + m[i] = make([]E, n) + } + return m +} diff --git a/internal/poseidon2/matrices_test.go b/internal/poseidon2/matrices_test.go new file mode 100644 index 0000000000..adb75dd840 --- /dev/null +++ b/internal/poseidon2/matrices_test.go @@ -0,0 +1,130 @@ +// Copyright 2020-2026 Consensys Software Inc. +// Licensed under the Apache License, Version 2.0. See the LICENSE file for details. + +package poseidon2 + +import ( + "testing" + + fr "github.com/consensys/gnark-crypto/field/koalabear" +) + +func fromRows(rows [][]int64) [][]fr.Element { + m := newMatrix[fr.Element](len(rows)) + for i := range rows { + for j := range rows[i] { + m[i][j].SetInt64(rows[i][j]) + } + } + return m +} + +func requireEqual(t *testing.T, name string, got, want [][]fr.Element) { + t.Helper() + if len(got) != len(want) { + t.Fatalf("%s: got %d rows, want %d", name, len(got), len(want)) + } + for i := range want { + if len(got[i]) != len(want[i]) { + t.Fatalf("%s: row %d has %d entries, want %d", name, i, len(got[i]), len(want[i])) + } + for j := range want[i] { + if !got[i][j].Equal(&want[i][j]) { + t.Fatalf("%s: entry (%d, %d) is %s, want %s", name, i, j, got[i][j].String(), want[i][j].String()) + } + } + } +} + +func TestExternalMatrixSmallWidths(t *testing.T) { + for _, kind := range []ExternalMatrixKind{Plonky3, Paper} { + requireEqual(t, "width 2", ExternalMatrix[fr.Element](2, kind), fromRows([][]int64{ + {2, 1}, + {1, 2}, + })) + requireEqual(t, "width 3", ExternalMatrix[fr.Element](3, kind), fromRows([][]int64{ + {2, 1, 1}, + {1, 2, 1}, + {1, 1, 2}, + })) + } +} + +func TestExternalMatrixWidth4IsM4(t *testing.T) { + requireEqual(t, "plonky3", ExternalMatrix[fr.Element](4, Plonky3), fromRows([][]int64{ + {2, 3, 1, 1}, + {1, 2, 3, 1}, + {1, 1, 2, 3}, + {3, 1, 1, 2}, + })) + requireEqual(t, "paper", ExternalMatrix[fr.Element](4, Paper), fromRows([][]int64{ + {5, 7, 1, 3}, + {4, 6, 1, 1}, + {1, 3, 5, 7}, + {1, 1, 4, 6}, + })) +} + +// TestExternalMatrixCirculant checks circ(2·M4, M4, ..., M4) for the widths 4k, k > 1. +func TestExternalMatrixCirculant(t *testing.T) { + for _, kind := range []ExternalMatrixKind{Plonky3, Paper} { + for _, width := range []int{8, 12, 16, 24} { + got := ExternalMatrix[fr.Element](width, kind) + want := newMatrix[fr.Element](width) + for bi := range width / 4 { + for bj := range width / 4 { + for r := range 4 { + for c := range 4 { + v := m4Blocks[kind][r][c] + if bi == bj { + v *= 2 + } + want[4*bi+r][4*bj+c].SetInt64(v) + } + } + } + } + requireEqual(t, "circulant", got, want) + } + } +} + +func TestExternalMatrixPanics(t *testing.T) { + for _, width := range []int{-4, 0, 1, 5, 6, 7, 10} { + func() { + defer func() { + if recover() == nil { + t.Fatalf("width %d: expected a panic", width) + } + }() + ExternalMatrix[fr.Element](width, Plonky3) + }() + } + func() { + defer func() { + if recover() == nil { + t.Fatal("expected a panic for an unknown kind") + } + }() + ExternalMatrix[fr.Element](4, ExternalMatrixKind(7)) + }() +} + +func TestInternalMatrix(t *testing.T) { + d := make([]fr.Element, 4) + for i, v := range []int64{3, -1, 0, 7} { + d[i].SetInt64(v) + } + requireEqual(t, "internal", InternalMatrix(d), fromRows([][]int64{ + {4, 1, 1, 1}, + {1, 0, 1, 1}, + {1, 1, 1, 1}, + {1, 1, 1, 8}, + })) + // d is not modified + var want fr.Element + want.SetInt64(3) + if !d[0].Equal(&want) { + t.Fatal("InternalMatrix modified its argument") + } +} From 260af680fa6132417000120c7a4ea508277fa256 Mon Sep 17 00:00:00 2001 From: Arya Tabaie Date: Tue, 6 Oct 2026 15:17:50 -0500 Subject: [PATCH 13/18] build: go generate Signed-off-by: Arya Tabaie --- field/mamabear/poseidon2/poseidon2_test.go | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/field/mamabear/poseidon2/poseidon2_test.go b/field/mamabear/poseidon2/poseidon2_test.go index 54d50dad12..fd59e1494c 100644 --- a/field/mamabear/poseidon2/poseidon2_test.go +++ b/field/mamabear/poseidon2/poseidon2_test.go @@ -80,8 +80,8 @@ func TestDenseMatricesMatchInPlace(t *testing.T) { for _, tc := range []struct { width, nbFullRounds, nbPartialRounds int }{ - {16, 6, 21}, - {24, 6, 21}, + {16, 8, 32}, + {24, 8, 32}, } { h := NewPermutation(tc.width, tc.nbFullRounds, tc.nbPartialRounds) From 11d808ebc133b8556f5051ac1c963b133636af8a Mon Sep 17 00:00:00 2001 From: Arya Tabaie Date: Tue, 6 Oct 2026 16:34:01 -0500 Subject: [PATCH 14/18] feat: polynomials for mamabear.E3 Signed-off-by: Arya Tabaie --- field/mamabear/extensions/e3_vector.go | 19 ++ field/mamabear/extensions/polynomial/doc.go | 7 + .../extensions/polynomial/multilin_e3.go | 242 +++++++++++++++++ .../extensions/polynomial/multilin_e3_test.go | 148 ++++++++++ .../extensions/polynomial/polynomial_e3.go | 253 ++++++++++++++++++ .../polynomial/polynomial_e3_test.go | 245 +++++++++++++++++ .../mamabear/extensions/polynomial/pool_e3.go | 191 +++++++++++++ internal/generator/config/fields.go | 1 + .../template/extensions/e3vector.go.tmpl | 19 ++ 9 files changed, 1125 insertions(+) create mode 100644 field/mamabear/extensions/polynomial/doc.go create mode 100644 field/mamabear/extensions/polynomial/multilin_e3.go create mode 100644 field/mamabear/extensions/polynomial/multilin_e3_test.go create mode 100644 field/mamabear/extensions/polynomial/polynomial_e3.go create mode 100644 field/mamabear/extensions/polynomial/polynomial_e3_test.go create mode 100644 field/mamabear/extensions/polynomial/pool_e3.go diff --git a/field/mamabear/extensions/e3_vector.go b/field/mamabear/extensions/e3_vector.go index 51a53ef9e0..ebe602c18d 100644 --- a/field/mamabear/extensions/e3_vector.go +++ b/field/mamabear/extensions/e3_vector.go @@ -179,3 +179,22 @@ func mulAccByElementGeneric(vector VectorE3, scale []fr.Element, alpha *E3) { vector[i].Add(&vector[i], &tmp) } } + +// SetRandom sets all elements of vector to random values, returning the first error encountered, if any. +func (vector VectorE3) SetRandom() error { + for i := range vector { + if _, err := vector[i].SetRandom(); err != nil { + return err + } + } + return nil +} + +// MustSetRandom sets all elements of vector to random values, panicking if an error is encountered. +func (vector VectorE3) MustSetRandom() { + for i := range vector { + if _, err := vector[i].SetRandom(); err != nil { + panic(err) + } + } +} diff --git a/field/mamabear/extensions/polynomial/doc.go b/field/mamabear/extensions/polynomial/doc.go new file mode 100644 index 0000000000..aa346f3ea3 --- /dev/null +++ b/field/mamabear/extensions/polynomial/doc.go @@ -0,0 +1,7 @@ +// Copyright 2020-2026 Consensys Software Inc. +// Licensed under the Apache License, Version 2.0. See the LICENSE file for details. + +// Code generated by consensys/gnark-crypto DO NOT EDIT + +// Package polynomial provides polynomial methods and commitment schemes. +package polynomial diff --git a/field/mamabear/extensions/polynomial/multilin_e3.go b/field/mamabear/extensions/polynomial/multilin_e3.go new file mode 100644 index 0000000000..4a4fba896d --- /dev/null +++ b/field/mamabear/extensions/polynomial/multilin_e3.go @@ -0,0 +1,242 @@ +// Copyright 2020-2026 Consensys Software Inc. +// Licensed under the Apache License, Version 2.0. See the LICENSE file for details. + +// Code generated by consensys/gnark-crypto DO NOT EDIT + +package polynomial + +import ( + "math/bits" + + fr "github.com/consensys/gnark-crypto/field/mamabear" + "github.com/consensys/gnark-crypto/field/mamabear/extensions" + basepoly "github.com/consensys/gnark-crypto/field/mamabear/polynomial" + "github.com/consensys/gnark-crypto/utils" +) + +// MultiLinE3 tracks the values of a (dense i.e. not sparse) multilinear polynomial +// The variables are X₁ through Xₙ where n = log(len(.)) +// .[∑ᵢ 2ⁱ⁻¹ bₙ₋ᵢ] = the polynomial evaluated at (b₁, b₂, ..., bₙ) +// It is understood that any hypercube evaluation can be extrapolated to a multilinear polynomial +type MultiLinE3 []extensions.E3 + +// Fold is partial evaluation function k[X₁, X₂, ..., Xₙ] → k[X₂, ..., Xₙ] by setting X₁=r +func (m *MultiLinE3) Fold(r extensions.E3) { + mid := len(*m) / 2 + + bottom, top := (*m)[:mid], (*m)[mid:] + + var t extensions.E3 // no need to update the top part + + // updating bookkeeping table + // knowing that the polynomial f ∈ (k[X₂, ..., Xₙ])[X₁] is linear, we would get f(r) = f(0) + r(f(1) - f(0)) + // the following loop computes the evaluations of f(r) accordingly: + // f(r, b₂, ..., bₙ) = f(0, b₂, ..., bₙ) + r(f(1, b₂, ..., bₙ) - f(0, b₂, ..., bₙ)) + for i := range mid { + // table[i] ← table[i] + r (table[i + mid] - table[i]) + t.Sub(&top[i], &bottom[i]) + t.Mul(&t, &r) + bottom[i].Add(&bottom[i], &t) + } + + *m = (*m)[:mid] +} + +func (m *MultiLinE3) FoldParallel(r extensions.E3) utils.Task { + mid := len(*m) / 2 + bottom, top := (*m)[:mid], (*m)[mid:] + + *m = bottom + + return func(start, end int) { + var t extensions.E3 // no need to update the top part + for i := start; i < end; i++ { + // table[i] ← table[i] + r (table[i + mid] - table[i]) + t.Sub(&top[i], &bottom[i]) + t.Mul(&t, &r) + bottom[i].Add(&bottom[i], &t) + } + } +} + +// FoldFromBase sets m to the partial evaluation X₁=r of the multilinear polynomial whose +// hypercube evaluations are the base field elements b: +// +// m[i] = b[i] + r (b[i + len(b)/2] - b[i]) +// +// m's backing array is reused if it is large enough. +func (m *MultiLinE3) FoldFromBase(b basepoly.MultiLin, r *extensions.E3) { + mid := len(b) / 2 + m.resize(mid) + m.foldFromBase(b, r, 0, mid) +} + +// FoldFromBaseParallel is the parallel version of FoldFromBase. It sizes m, and returns +// a task that computes the entries m[start:end]. The task reads *r when it runs, so *r must +// not be modified until all tasks have completed. +func (m *MultiLinE3) FoldFromBaseParallel(b basepoly.MultiLin, r *extensions.E3) utils.Task { + mid := len(b) / 2 + m.resize(mid) + dst := *m + + return func(start, end int) { + dst.foldFromBase(b, r, start, end) + } +} + +// EvaluateBase evaluates, on the given coordinates, the multilinear polynomial whose hypercube +// evaluations are the base field elements b. It first folds b into m by the first coordinate +// (see FoldFromBase), and then folds m by the remaining coordinates. m serves as scratch space, +// is overwritten, and its backing array is reused if it is large enough. +// Unlike Evaluate, no copy of the table is made, and b is left untouched. +func (m *MultiLinE3) EvaluateBase(b basepoly.MultiLin, coordinates []extensions.E3) extensions.E3 { + if len(coordinates) == 0 { + m.resize(1) + return *(*m)[0].SetElement(&b[0]) + } + + m.FoldFromBase(b, &coordinates[0]) + for _, r := range coordinates[1:] { + m.Fold(r) + } + return (*m)[0] +} + +func (m *MultiLinE3) resize(n int) { + if cap(*m) >= n { + *m = (*m)[:n] + } else { + *m = make(MultiLinE3, n) + } +} + +func (m MultiLinE3) foldFromBase(b basepoly.MultiLin, r *extensions.E3, start, end int) { + mid := len(b) / 2 + var diff fr.Element + for i := start; i < end; i++ { + diff.Sub(&b[mid+i], &b[i]) + m[i].MulByElement(r, &diff) + m[i].AddElement(&m[i], &b[i]) + } +} + +func (m MultiLinE3) Sum() extensions.E3 { + s := m[0] + for i := 1; i < len(m); i++ { + s.Add(&s, &m[i]) + } + return s +} + +func _cloneE3(m MultiLinE3, p *PoolE3) MultiLinE3 { + if p == nil { + return m.Clone() + } else { + return p.Clone(m) + } +} + +func _dumpE3(m MultiLinE3, p *PoolE3) { + if p != nil { + p.Dump(m) + } +} + +// Evaluate extrapolate the value of the multilinear polynomial corresponding to m +// on the given coordinates +func (m MultiLinE3) Evaluate(coordinates []extensions.E3, p *PoolE3) extensions.E3 { + // Folding is a mutating operation + bkCopy := _cloneE3(m, p) + + // Evaluate step by step through repeated folding (i.e. evaluation at the first remaining variable) + for _, r := range coordinates { + bkCopy.Fold(r) + } + + result := bkCopy[0] + + _dumpE3(bkCopy, p) + return result +} + +// Clone creates a deep copy of a bookkeeping table. +// Both multilinear interpolation and sumcheck require folding an underlying +// array, but folding changes the array. To do both one requires a deep copy +// of the bookkeeping table. +func (m MultiLinE3) Clone() MultiLinE3 { + res := make(MultiLinE3, len(m)) + copy(res, m) + return res +} + +// Add two bookKeepingTables +func (m *MultiLinE3) Add(left, right MultiLinE3) { + size := len(left) + // Check that left and right have the same size + if len(right) != size || len(*m) != size { + panic("left, right and destination must have the right size") + } + + // Add elementwise + for i := range size { + (*m)[i].Add(&left[i], &right[i]) + } +} + +// EvalEqE3 computes Eq(q₁, ... , qₙ, h₁, ... , hₙ) = Π₁ⁿ Eq(qᵢ, hᵢ) +// where Eq(x,y) = xy + (1-x)(1-y) = 1 - x - y + xy + xy interpolates +// +// _________________ +// | | | +// | 0 | 1 | +// |_______|_______| +// y | | | +// | 1 | 0 | +// |_______|_______| +// +// x +// +// In other words the polynomial evaluated here is the multilinear extrapolation of +// one that evaluates to q' == h' for vectors q', h' of binary values +func EvalEqE3(q, h []extensions.E3) extensions.E3 { + var res, nxt, one, sum extensions.E3 + one.SetOne() + for i := range len(q) { + nxt.Mul(&q[i], &h[i]) // nxt <- qᵢ * hᵢ + nxt.Double(&nxt) // nxt <- 2 * qᵢ * hᵢ + nxt.Add(&nxt, &one) // nxt <- 1 + 2 * qᵢ * hᵢ + sum.Add(&q[i], &h[i]) // sum <- qᵢ + hᵢ TODO: Why not subtract one by one from nxt? More parallel? + + if i == 0 { + res.Sub(&nxt, &sum) // nxt <- 1 + 2 * qᵢ * hᵢ - qᵢ - hᵢ + } else { + nxt.Sub(&nxt, &sum) // nxt <- 1 + 2 * qᵢ * hᵢ - qᵢ - hᵢ + res.Mul(&res, &nxt) // res <- res * nxt + } + } + return res +} + +// Eq sets m to the representation of the polynomial Eq(q₁, ..., qₙ, *, ..., *) × m[0] +func (m *MultiLinE3) Eq(q []extensions.E3) { + n := len(q) + + if len(*m) != 1<= 0; i-- { + res.Mul(&res, v) + res.Add(&res, &(*p)[i]) + } + + return res +} + +// Clone returns a copy of the polynomial +func (p *PolynomialE3) Clone() PolynomialE3 { + _p := make(PolynomialE3, len(*p)) + copy(_p, *p) + return _p +} + +// Set to another polynomial +func (p *PolynomialE3) Set(p1 PolynomialE3) { + if len(*p) != len(p1) { + *p = p1.Clone() + return + } + + for i := range len(p1) { + (*p)[i].Set(&p1[i]) + } +} + +// AddConstantInPlace adds a constant to the polynomial, modifying p +func (p *PolynomialE3) AddConstantInPlace(c *extensions.E3) { + for i := range len(*p) { + (*p)[i].Add(&(*p)[i], c) + } +} + +// SubConstantInPlace subs a constant to the polynomial, modifying p +func (p *PolynomialE3) SubConstantInPlace(c *extensions.E3) { + for i := range len(*p) { + (*p)[i].Sub(&(*p)[i], c) + } +} + +// ScaleInPlace multiplies p by v, modifying p +func (p *PolynomialE3) ScaleInPlace(c *extensions.E3) { + for i := range len(*p) { + (*p)[i].Mul(&(*p)[i], c) + } +} + +// Scale multiplies p0 by v, storing the result in p +func (p *PolynomialE3) Scale(c *extensions.E3, p0 PolynomialE3) { + if len(*p) != len(p0) { + *p = make(PolynomialE3, len(p0)) + } + for i := range len(p0) { + (*p)[i].Mul(c, &p0[i]) + } +} + +// Add adds p1 to p2 +// This function allocates a new slice unless p == p1 or p == p2 +func (p *PolynomialE3) Add(p1, p2 PolynomialE3) *PolynomialE3 { + + bigger := p1 + smaller := p2 + if len(bigger) < len(smaller) { + bigger, smaller = smaller, bigger + } + + if len(*p) == len(bigger) && (&(*p)[0] == &bigger[0]) { + for i := range len(smaller) { + (*p)[i].Add(&(*p)[i], &smaller[i]) + } + return p + } + + if len(*p) == len(smaller) && (&(*p)[0] == &smaller[0]) { + for i := range len(smaller) { + (*p)[i].Add(&(*p)[i], &bigger[i]) + } + *p = append(*p, bigger[len(smaller):]...) + return p + } + + res := make(PolynomialE3, len(bigger)) + copy(res, bigger) + for i := range len(smaller) { + res[i].Add(&res[i], &smaller[i]) + } + *p = res + return p +} + +// Sub subtracts p2 from p1 +// TODO make interface more consistent with Add +func (p *PolynomialE3) Sub(p1, p2 PolynomialE3) *PolynomialE3 { + if len(p1) != len(p2) || len(p2) != len(*p) { + return nil + } + for i := range len(*p) { + (*p)[i].Sub(&p1[i], &p2[i]) + } + return p +} + +// Equal checks equality between two polynomials +func (p *PolynomialE3) Equal(p1 PolynomialE3) bool { + if (*p == nil) != (p1 == nil) { + return false + } + + if len(*p) != len(p1) { + return false + } + + for i := range p1 { + if !(*p)[i].Equal(&p1[i]) { + return false + } + } + + return true +} + +func (p PolynomialE3) SetZero() { + for i := range len(p) { + p[i].SetZero() + } +} + +// InterpolateOnRangeE3 maps vector v to polynomial f +// such that f(i) = v[i] for 0 ≤ i < len(v). +// len(f) = len(v) and deg(f) ≤ len(v) - 1 +func InterpolateOnRangeE3(v []extensions.E3) PolynomialE3 { + nEvals := uint8(len(v)) + if int(nEvals) != len(v) { + panic("interpolation method too inefficient for nEvals > 255") + } + lagrange := getLagrangeBasisE3(nEvals) + + var res PolynomialE3 + res.Scale(&v[0], lagrange[0]) + + temp := make(PolynomialE3, nEvals) + + for i := uint8(1); i < nEvals; i++ { + temp.Scale(&v[i], lagrange[i]) + res.Add(res, temp) + } + + return res +} + +// lagrange bases used by InterpolateOnRangeE3 +var lagrangeBasisE3 sync.Map + +func getLagrangeBasisE3(domainSize uint8) []PolynomialE3 { + if res, ok := lagrangeBasisE3.Load(domainSize); ok { + return res.([]PolynomialE3) + } + + // not found. compute + var res []PolynomialE3 + if domainSize >= 2 { + res = computeLagrangeBasisE3(domainSize) + } else if domainSize == 1 { + res = []PolynomialE3{make(PolynomialE3, 1)} + res[0][0].SetOne() + } + lagrangeBasisE3.Store(domainSize, res) + + return res +} + +// computeLagrangeBasisE3 precomputes in explicit coefficient form for each 0 ≤ l < domainSize the polynomial +// pₗ := X (X-1) ... (X-l-1) (X-l+1) ... (X - domainSize + 1) / ( l (l-1) ... 2 (-1) ... (l - domainSize +1) ) +// Note that pₗ(l) = 1 and pₗ(n) = 0 if 0 ≤ l < domainSize, n ≠ l +func computeLagrangeBasisE3(domainSize uint8) []PolynomialE3 { + + constTerms := make([]extensions.E3, domainSize) + for i := range domainSize { + constTerms[i].SetInt64(-int64(i)) + } + + res := make([]PolynomialE3, domainSize) + multScratch := make(PolynomialE3, domainSize-1) + + // compute pₗ + for l := range domainSize { + + // TODO @Tabaie Optimize this with some trees? O(log(domainSize)) polynomial mults instead of O(domainSize)? Then again it would be fewer big poly mults vs many small poly mults + d := uint8(0) //d is the current degree of res + for i := range domainSize { + if i == l { + continue + } + if d == 0 { + res[l] = make(PolynomialE3, domainSize) + res[l][domainSize-2] = constTerms[i] + res[l][domainSize-1].SetOne() + } else { + current := res[l][domainSize-d-2:] + timesConst := multScratch[domainSize-d-2:] + + timesConst.Scale(&constTerms[i], current[1:]) //TODO: Directly double and add since constTerms are tiny? (even less than 4 bits) + nonLeading := current[0 : d+1] + + nonLeading.Add(nonLeading, timesConst) + + } + d++ + } + + } + + // We have pₗ(i≠l)=0. Now scale so that pₗ(l)=1 + // Replace the constTerms with norms + for l := range domainSize { + constTerms[l].Neg(&constTerms[l]) + constTerms[l] = res[l].Eval(&constTerms[l]) + } + constTerms = extensions.BatchInvertE3(constTerms) + for l := range domainSize { + res[l].ScaleInPlace(&constTerms[l]) + } + + return res +} diff --git a/field/mamabear/extensions/polynomial/polynomial_e3_test.go b/field/mamabear/extensions/polynomial/polynomial_e3_test.go new file mode 100644 index 0000000000..a9b6db0bda --- /dev/null +++ b/field/mamabear/extensions/polynomial/polynomial_e3_test.go @@ -0,0 +1,245 @@ +// Copyright 2020-2026 Consensys Software Inc. +// Licensed under the Apache License, Version 2.0. See the LICENSE file for details. + +// Code generated by consensys/gnark-crypto DO NOT EDIT + +package polynomial + +import ( + "math/big" + "testing" + + "github.com/consensys/gnark-crypto/field/mamabear/extensions" + "github.com/leanovate/gopter" + "github.com/leanovate/gopter/gen" + "github.com/leanovate/gopter/prop" + "github.com/stretchr/testify/assert" +) + +func TestPolynomialEvalE3(t *testing.T) { + + // build polynomial + f := make(PolynomialE3, 20) + for i := range 20 { + f[i].SetOne() + } + + // random value + var point extensions.E3 + point.MustSetRandom() + + // compute manually f(val) + var expectedEval, one, den extensions.E3 + var expo big.Int + one.SetOne() + expo.SetUint64(20) + expectedEval.Exp(point, &expo). + Sub(&expectedEval, &one) + den.Sub(&point, &one) + expectedEval.Div(&expectedEval, &den) + + // compute purported evaluation + purportedEval := f.Eval(&point) + + // check + if !purportedEval.Equal(&expectedEval) { + t.Fatal("polynomial evaluation failed") + } +} + +func TestPolynomialAddConstantInPlaceE3(t *testing.T) { + + // build polynomial + f := make(PolynomialE3, 20) + for i := range 20 { + f[i].SetOne() + } + + // constant to add + var c extensions.E3 + c.MustSetRandom() + + // add constant + f.AddConstantInPlace(&c) + + // check + var expectedCoeffs, one extensions.E3 + one.SetOne() + expectedCoeffs.Add(&one, &c) + for i := range 20 { + if !f[i].Equal(&expectedCoeffs) { + t.Fatal("AddConstantInPlace failed") + } + } +} + +func TestPolynomialSubConstantInPlaceE3(t *testing.T) { + + // build polynomial + f := make(PolynomialE3, 20) + for i := range 20 { + f[i].SetOne() + } + + // constant to sub + var c extensions.E3 + c.MustSetRandom() + + // sub constant + f.SubConstantInPlace(&c) + + // check + var expectedCoeffs, one extensions.E3 + one.SetOne() + expectedCoeffs.Sub(&one, &c) + for i := range 20 { + if !f[i].Equal(&expectedCoeffs) { + t.Fatal("SubConstantInPlace failed") + } + } +} + +func TestPolynomialScaleInPlaceE3(t *testing.T) { + + // build polynomial + f := make(PolynomialE3, 20) + for i := range 20 { + f[i].SetOne() + } + + // constant to scale by + var c extensions.E3 + c.MustSetRandom() + + // scale by constant + f.ScaleInPlace(&c) + + // check + for i := range 20 { + if !f[i].Equal(&c) { + t.Fatal("ScaleInPlace failed") + } + } + +} + +func TestPolynomialAddE3(t *testing.T) { + + // build unbalanced polynomials + f1 := make(PolynomialE3, 20) + f1Backup := make(PolynomialE3, 20) + for i := range 20 { + f1[i].SetOne() + f1Backup[i].SetOne() + } + f2 := make(PolynomialE3, 10) + f2Backup := make(PolynomialE3, 10) + for i := range 10 { + f2[i].SetOne() + f2Backup[i].SetOne() + } + + // expected result + var one, two extensions.E3 + one.SetOne() + two.Double(&one) + expectedSum := make(PolynomialE3, 20) + for i := range 10 { + expectedSum[i].Set(&two) + } + for i := 10; i < 20; i++ { + expectedSum[i].Set(&one) + } + + // caller is empty + var g PolynomialE3 + g.Add(f1, f2) + if !g.Equal(expectedSum) { + t.Fatal("add polynomials fails") + } + if !f1.Equal(f1Backup) { + t.Fatal("side effect, f1 should not have been modified") + } + if !f2.Equal(f2Backup) { + t.Fatal("side effect, f2 should not have been modified") + } + + // all operands are distinct + _f1 := f1.Clone() + _f1.Add(f1, f2) + if !_f1.Equal(expectedSum) { + t.Fatal("add polynomials fails") + } + if !f1.Equal(f1Backup) { + t.Fatal("side effect, f1 should not have been modified") + } + if !f2.Equal(f2Backup) { + t.Fatal("side effect, f2 should not have been modified") + } + + // first operand = caller + _f1 = f1.Clone() + _f2 := f2.Clone() + _f1.Add(_f1, _f2) + if !_f1.Equal(expectedSum) { + t.Fatal("add polynomials fails") + } + if !_f2.Equal(f2Backup) { + t.Fatal("side effect, _f2 should not have been modified") + } + + // second operand = caller + _f1 = f1.Clone() + _f2 = f2.Clone() + _f1.Add(_f2, _f1) + if !_f1.Equal(expectedSum) { + t.Fatal("add polynomials fails") + } + if !_f2.Equal(f2Backup) { + t.Fatal("side effect, _f2 should not have been modified") + } +} + +func TestPrecomputeLagrangeE3(t *testing.T) { + + testForDomainSize := func(domainSize uint8) bool { + polys := computeLagrangeBasisE3(domainSize) + + for l := range domainSize { + for i := range domainSize { + var I extensions.E3 + I.SetUint64(uint64(i)) + y := polys[l].Eval(&I) + + if i == l && !y.IsOne() || i != l && !y.IsZero() { + t.Errorf("domainSize = %d: p_%d(%d) = %s", domainSize, l, i, y.String()) + return false + } + } + } + return true + } + + t.Parallel() + parameters := gopter.DefaultTestParameters() + + const maxLagrangeDomainSize = 12 + + parameters.MinSuccessfulTests = maxLagrangeDomainSize + + properties := gopter.NewProperties(parameters) + + properties.Property("l'th lagrange polynomials must evaluate to 1 on l and 0 on other values in the domain", prop.ForAll( + testForDomainSize, + gen.UInt8Range(2, maxLagrangeDomainSize), + )) + + properties.TestingRun(t, gopter.ConsoleReporter(false)) +} + +func TestLagrangeCacheE3(t *testing.T) { + for _, i := range []int{5, 2, 8, 4, 6, 3, 0} { + b := getLagrangeBasisE3(uint8(i)) + assert.Equal(t, b, getLagrangeBasisE3(uint8(i))) // second call must yield the same result + } +} diff --git a/field/mamabear/extensions/polynomial/pool_e3.go b/field/mamabear/extensions/polynomial/pool_e3.go new file mode 100644 index 0000000000..e18c9faa45 --- /dev/null +++ b/field/mamabear/extensions/polynomial/pool_e3.go @@ -0,0 +1,191 @@ +// Copyright 2020-2026 Consensys Software Inc. +// Licensed under the Apache License, Version 2.0. See the LICENSE file for details. + +// Code generated by consensys/gnark-crypto DO NOT EDIT + +package polynomial + +import ( + "encoding/json" + "fmt" + "runtime" + "sort" + "sync" + "unsafe" + + "github.com/consensys/gnark-crypto/field/mamabear/extensions" +) + +// Memory management for polynomials +// WARNING: This is not thread safe TODO: Make sure that is not a problem +// TODO: There is a lot of "unsafe" memory management here and needs to be vetted thoroughly + +type sizedPoolE3 struct { + maxN int + pool sync.Pool + stats poolStatsE3 +} + +type inUseDataE3 struct { + allocatedFor []uintptr + pool *sizedPoolE3 +} + +type PoolE3 struct { + //lock sync.Mutex + inUse sync.Map + subPools []sizedPoolE3 +} + +func (p *sizedPoolE3) get(n int) *extensions.E3 { + p.stats.make(n) + return p.pool.Get().(*extensions.E3) +} + +func (p *sizedPoolE3) put(ptr *extensions.E3) { + p.stats.dump() + p.pool.Put(ptr) +} + +func NewPoolE3(maxN ...int) (pool PoolE3) { + + sort.Ints(maxN) + pool = PoolE3{ + subPools: make([]sizedPoolE3, len(maxN)), + } + + for i := range pool.subPools { + subPool := &pool.subPools[i] + subPool.maxN = maxN[i] + subPool.pool = sync.Pool{ + New: func() any { + subPool.stats.Allocated++ + return getDataPointerE3(make([]extensions.E3, 0, subPool.maxN)) + }, + } + } + return +} + +func (p *PoolE3) findCorrespondingPool(n int) *sizedPoolE3 { + poolI := 0 + for poolI < len(p.subPools) && n > p.subPools[poolI].maxN { + poolI++ + } + return &p.subPools[poolI] // out of bounds error here would mean that n is too large +} + +func (p *PoolE3) Make(n int) []extensions.E3 { + pool := p.findCorrespondingPool(n) + ptr := pool.get(n) + p.addInUse(ptr, pool) + return unsafe.Slice(ptr, n) +} + +// Dump dumps a set of polynomials into the pool +func (p *PoolE3) Dump(slices ...[]extensions.E3) { + for _, slice := range slices { + ptr := getDataPointerE3(slice) + if metadata, ok := p.inUse.Load(ptr); ok { + p.inUse.Delete(ptr) + metadata.(inUseDataE3).pool.put(ptr) + } else { + panic("attempting to dump a slice not created by the pool") + } + } +} + +func (p *PoolE3) addInUse(ptr *extensions.E3, pool *sizedPoolE3) { + pcs := make([]uintptr, 2) + n := runtime.Callers(3, pcs) + + if prevPcs, ok := p.inUse.Load(ptr); ok { // TODO: remove if unnecessary for security + panic(fmt.Errorf("re-allocated non-dumped slice, previously allocated at %v", runtime.CallersFrames(prevPcs.(inUseDataE3).allocatedFor))) + } + p.inUse.Store(ptr, inUseDataE3{ + allocatedFor: pcs[:n], + pool: pool, + }) +} + +func printFrameE3(frame runtime.Frame) { + fmt.Printf("\t%s line %d, function %s\n", frame.File, frame.Line, frame.Function) +} + +func (p *PoolE3) printInUse() { + fmt.Println("slices never dumped allocated at:") + p.inUse.Range(func(_, pcs any) bool { + fmt.Println("-------------------------") + + var frame runtime.Frame + frames := runtime.CallersFrames(pcs.(inUseDataE3).allocatedFor) + more := true + for more { + frame, more = frames.Next() + printFrameE3(frame) + } + return true + }) +} + +type poolStatsE3 struct { + Used int + Allocated int + ReuseRate float64 + InUse int + GreatestNUsed int + SmallestNUsed int +} + +type poolsStatsE3 struct { + SubPools []poolStatsE3 + InUse int +} + +func (s *poolStatsE3) make(n int) { + s.Used++ + s.InUse++ + if n > s.GreatestNUsed { + s.GreatestNUsed = n + } + if s.SmallestNUsed == 0 || s.SmallestNUsed > n { + s.SmallestNUsed = n + } +} + +func (s *poolStatsE3) dump() { + s.InUse-- +} + +func (s *poolStatsE3) finalize() { + s.ReuseRate = float64(s.Used) / float64(s.Allocated) +} + +func getDataPointerE3(slice []extensions.E3) *extensions.E3 { + return (*extensions.E3)(unsafe.SliceData(slice)) +} + +func (p *PoolE3) PrintPoolStats() { + InUse := 0 + subStats := make([]poolStatsE3, len(p.subPools)) + for i := range p.subPools { + subPool := &p.subPools[i] + subPool.stats.finalize() + subStats[i] = subPool.stats + InUse += subPool.stats.InUse + } + + stats := poolsStatsE3{ + SubPools: subStats, + InUse: InUse, + } + serialized, _ := json.MarshalIndent(stats, "", " ") + fmt.Println(string(serialized)) + p.printInUse() +} + +func (p *PoolE3) Clone(slice []extensions.E3) []extensions.E3 { + res := p.Make(len(slice)) + copy(res, slice) + return res +} diff --git a/internal/generator/config/fields.go b/internal/generator/config/fields.go index 23ba42032b..f55a93d9d7 100644 --- a/internal/generator/config/fields.go +++ b/internal/generator/config/fields.go @@ -46,5 +46,6 @@ func init() { // hand-written; the asm generator has no regime for a single-word // field with a sub-word radix. HandwrittenVectorASMAMD64: true, + PolynomialExtensions: []int{3}, }) } diff --git a/internal/generator/field/template/extensions/e3vector.go.tmpl b/internal/generator/field/template/extensions/e3vector.go.tmpl index 10a6c17b25..a9c543b119 100644 --- a/internal/generator/field/template/extensions/e3vector.go.tmpl +++ b/internal/generator/field/template/extensions/e3vector.go.tmpl @@ -172,3 +172,22 @@ func mulAccByElementGeneric(vector VectorE3, scale []fr.Element, alpha *E3) { vector[i].Add(&vector[i], &tmp) } } + +// SetRandom sets all elements of vector to random values, returning the first error encountered, if any. +func (vector VectorE3) SetRandom() error { + for i := range vector { + if _, err := vector[i].SetRandom(); err != nil { + return err + } + } + return nil +} + +// MustSetRandom sets all elements of vector to random values, panicking if an error is encountered. +func (vector VectorE3) MustSetRandom() { + for i := range vector { + if _, err := vector[i].SetRandom(); err != nil { + panic(err) + } + } +} From a9a36e8b8092851af9a15de9424541f6c4afdfc6 Mon Sep 17 00:00:00 2001 From: Arya Tabaie Date: Tue, 6 Oct 2026 16:44:39 -0500 Subject: [PATCH 15/18] feat: register mamabear poseidon2 Signed-off-by: Arya Tabaie --- hash/all/allhashes.go | 1 + 1 file changed, 1 insertion(+) diff --git a/hash/all/allhashes.go b/hash/all/allhashes.go index 6036a3c4cf..84b14658f2 100644 --- a/hash/all/allhashes.go +++ b/hash/all/allhashes.go @@ -25,6 +25,7 @@ import ( _ "github.com/consensys/gnark-crypto/field/babybear/poseidon2" _ "github.com/consensys/gnark-crypto/field/goldilocks/poseidon2" _ "github.com/consensys/gnark-crypto/field/koalabear/poseidon2" + _ "github.com/consensys/gnark-crypto/field/mamabear/poseidon2" _ "github.com/consensys/gnark-crypto/ecc/grumpkin/fr/mimc" _ "github.com/consensys/gnark-crypto/ecc/grumpkin/fr/poseidon2" From 54acd6fb1d084a18aa3270bc15d392e92d759cc3 Mon Sep 17 00:00:00 2001 From: Arya Tabaie Date: Wed, 7 Oct 2026 13:29:01 -0500 Subject: [PATCH 16/18] refactor: use MustSetRandom when possible. Signed-off-by: Arya Tabaie --- field/babybear/extensions/e6.go | 4 +--- field/koalabear/extensions/e6.go | 4 +--- .../extensions/polynomial/multilin_e6_test.go | 17 +++++------------ field/mamabear/extensions/e3_vector.go | 4 +--- .../extensions/polynomial/multilin_e3_test.go | 17 +++++------------ .../field/template/extensions/e3vector.go.tmpl | 4 +--- .../field/template/extensions/e6.go.tmpl | 4 +--- .../polynomial/template/multilin.test.go.tmpl | 16 ++++------------ 8 files changed, 19 insertions(+), 51 deletions(-) diff --git a/field/babybear/extensions/e6.go b/field/babybear/extensions/e6.go index 4d0167ffd9..daf7c608a5 100644 --- a/field/babybear/extensions/e6.go +++ b/field/babybear/extensions/e6.go @@ -604,9 +604,7 @@ func (vector VectorE6) SetRandom() error { // MustSetRandom sets all elements of vector to random values, panicking if an error is encountered. func (vector VectorE6) MustSetRandom() { for i := range vector { - if _, err := vector[i].SetRandom(); err != nil { - panic(err) - } + vector[i].MustSetRandom() } } diff --git a/field/koalabear/extensions/e6.go b/field/koalabear/extensions/e6.go index 342fb3e4e4..ad6dbee4bd 100644 --- a/field/koalabear/extensions/e6.go +++ b/field/koalabear/extensions/e6.go @@ -605,9 +605,7 @@ func (vector VectorE6) SetRandom() error { // MustSetRandom sets all elements of vector to random values, panicking if an error is encountered. func (vector VectorE6) MustSetRandom() { for i := range vector { - if _, err := vector[i].SetRandom(); err != nil { - panic(err) - } + vector[i].MustSetRandom() } } diff --git a/field/koalabear/extensions/polynomial/multilin_e6_test.go b/field/koalabear/extensions/polynomial/multilin_e6_test.go index e5a412f8a8..9b287b1818 100644 --- a/field/koalabear/extensions/polynomial/multilin_e6_test.go +++ b/field/koalabear/extensions/polynomial/multilin_e6_test.go @@ -8,6 +8,7 @@ package polynomial import ( "testing" + fr "github.com/consensys/gnark-crypto/field/koalabear" "github.com/consensys/gnark-crypto/field/koalabear/extensions" basepoly "github.com/consensys/gnark-crypto/field/koalabear/polynomial" "github.com/stretchr/testify/assert" @@ -88,9 +89,7 @@ func TestFoldedEqTableE6(t *testing.T) { func TestFoldFromBaseE6(t *testing.T) { for _, n := range []int{2, 4, 8, 64} { b := make(basepoly.MultiLin, n) - for i := range b { - b[i].MustSetRandom() - } + fr.Vector(b).MustSetRandom() var r extensions.E6 r.MustSetRandom() @@ -118,13 +117,9 @@ func TestFoldFromBaseE6(t *testing.T) { func TestEvaluateBaseE6(t *testing.T) { for _, nbVars := range []int{0, 1, 2, 3, 6} { b := make(basepoly.MultiLin, 1< Date: Wed, 7 Oct 2026 13:33:28 -0500 Subject: [PATCH 17/18] refactor: only generate polynomial package for base if doing so for at least one ext Signed-off-by: Arya Tabaie --- field/babybear/polynomial/doc.go | 7 - field/babybear/polynomial/multilin.go | 179 ---------- field/babybear/polynomial/multilin_test.go | 85 ----- field/babybear/polynomial/polynomial.go | 310 ------------------ field/babybear/polynomial/polynomial_test.go | 255 -------------- field/babybear/polynomial/pool.go | 191 ----------- field/goldilocks/polynomial/doc.go | 7 - field/goldilocks/polynomial/multilin.go | 179 ---------- field/goldilocks/polynomial/multilin_test.go | 85 ----- field/goldilocks/polynomial/polynomial.go | 310 ------------------ .../goldilocks/polynomial/polynomial_test.go | 255 -------------- field/goldilocks/polynomial/pool.go | 191 ----------- internal/generator/main.go | 16 +- 13 files changed, 10 insertions(+), 2060 deletions(-) delete mode 100644 field/babybear/polynomial/doc.go delete mode 100644 field/babybear/polynomial/multilin.go delete mode 100644 field/babybear/polynomial/multilin_test.go delete mode 100644 field/babybear/polynomial/polynomial.go delete mode 100644 field/babybear/polynomial/polynomial_test.go delete mode 100644 field/babybear/polynomial/pool.go delete mode 100644 field/goldilocks/polynomial/doc.go delete mode 100644 field/goldilocks/polynomial/multilin.go delete mode 100644 field/goldilocks/polynomial/multilin_test.go delete mode 100644 field/goldilocks/polynomial/polynomial.go delete mode 100644 field/goldilocks/polynomial/polynomial_test.go delete mode 100644 field/goldilocks/polynomial/pool.go diff --git a/field/babybear/polynomial/doc.go b/field/babybear/polynomial/doc.go deleted file mode 100644 index aa346f3ea3..0000000000 --- a/field/babybear/polynomial/doc.go +++ /dev/null @@ -1,7 +0,0 @@ -// Copyright 2020-2026 Consensys Software Inc. -// Licensed under the Apache License, Version 2.0. See the LICENSE file for details. - -// Code generated by consensys/gnark-crypto DO NOT EDIT - -// Package polynomial provides polynomial methods and commitment schemes. -package polynomial diff --git a/field/babybear/polynomial/multilin.go b/field/babybear/polynomial/multilin.go deleted file mode 100644 index f29a16a4fc..0000000000 --- a/field/babybear/polynomial/multilin.go +++ /dev/null @@ -1,179 +0,0 @@ -// Copyright 2020-2026 Consensys Software Inc. -// Licensed under the Apache License, Version 2.0. See the LICENSE file for details. - -// Code generated by consensys/gnark-crypto DO NOT EDIT - -package polynomial - -import ( - "math/bits" - - "github.com/consensys/gnark-crypto/field/babybear" - "github.com/consensys/gnark-crypto/utils" -) - -// MultiLin tracks the values of a (dense i.e. not sparse) multilinear polynomial -// The variables are X₁ through Xₙ where n = log(len(.)) -// .[∑ᵢ 2ⁱ⁻¹ bₙ₋ᵢ] = the polynomial evaluated at (b₁, b₂, ..., bₙ) -// It is understood that any hypercube evaluation can be extrapolated to a multilinear polynomial -type MultiLin []babybear.Element - -// Fold is partial evaluation function k[X₁, X₂, ..., Xₙ] → k[X₂, ..., Xₙ] by setting X₁=r -func (m *MultiLin) Fold(r babybear.Element) { - mid := len(*m) / 2 - - bottom, top := (*m)[:mid], (*m)[mid:] - - var t babybear.Element // no need to update the top part - - // updating bookkeeping table - // knowing that the polynomial f ∈ (k[X₂, ..., Xₙ])[X₁] is linear, we would get f(r) = f(0) + r(f(1) - f(0)) - // the following loop computes the evaluations of f(r) accordingly: - // f(r, b₂, ..., bₙ) = f(0, b₂, ..., bₙ) + r(f(1, b₂, ..., bₙ) - f(0, b₂, ..., bₙ)) - for i := range mid { - // table[i] ← table[i] + r (table[i + mid] - table[i]) - t.Sub(&top[i], &bottom[i]) - t.Mul(&t, &r) - bottom[i].Add(&bottom[i], &t) - } - - *m = (*m)[:mid] -} - -func (m *MultiLin) FoldParallel(r babybear.Element) utils.Task { - mid := len(*m) / 2 - bottom, top := (*m)[:mid], (*m)[mid:] - - *m = bottom - - return func(start, end int) { - var t babybear.Element // no need to update the top part - for i := start; i < end; i++ { - // table[i] ← table[i] + r (table[i + mid] - table[i]) - t.Sub(&top[i], &bottom[i]) - t.Mul(&t, &r) - bottom[i].Add(&bottom[i], &t) - } - } -} - -func (m MultiLin) Sum() babybear.Element { - s := m[0] - for i := 1; i < len(m); i++ { - s.Add(&s, &m[i]) - } - return s -} - -func _clone(m MultiLin, p *Pool) MultiLin { - if p == nil { - return m.Clone() - } else { - return p.Clone(m) - } -} - -func _dump(m MultiLin, p *Pool) { - if p != nil { - p.Dump(m) - } -} - -// Evaluate extrapolate the value of the multilinear polynomial corresponding to m -// on the given coordinates -func (m MultiLin) Evaluate(coordinates []babybear.Element, p *Pool) babybear.Element { - // Folding is a mutating operation - bkCopy := _clone(m, p) - - // Evaluate step by step through repeated folding (i.e. evaluation at the first remaining variable) - for _, r := range coordinates { - bkCopy.Fold(r) - } - - result := bkCopy[0] - - _dump(bkCopy, p) - return result -} - -// Clone creates a deep copy of a bookkeeping table. -// Both multilinear interpolation and sumcheck require folding an underlying -// array, but folding changes the array. To do both one requires a deep copy -// of the bookkeeping table. -func (m MultiLin) Clone() MultiLin { - res := make(MultiLin, len(m)) - copy(res, m) - return res -} - -// Add two bookKeepingTables -func (m *MultiLin) Add(left, right MultiLin) { - size := len(left) - // Check that left and right have the same size - if len(right) != size || len(*m) != size { - panic("left, right and destination must have the right size") - } - - // Add elementwise - for i := range size { - (*m)[i].Add(&left[i], &right[i]) - } -} - -// EvalEq computes Eq(q₁, ... , qₙ, h₁, ... , hₙ) = Π₁ⁿ Eq(qᵢ, hᵢ) -// where Eq(x,y) = xy + (1-x)(1-y) = 1 - x - y + xy + xy interpolates -// -// _________________ -// | | | -// | 0 | 1 | -// |_______|_______| -// y | | | -// | 1 | 0 | -// |_______|_______| -// -// x -// -// In other words the polynomial evaluated here is the multilinear extrapolation of -// one that evaluates to q' == h' for vectors q', h' of binary values -func EvalEq(q, h []babybear.Element) babybear.Element { - var res, nxt, one, sum babybear.Element - one.SetOne() - for i := range len(q) { - nxt.Mul(&q[i], &h[i]) // nxt <- qᵢ * hᵢ - nxt.Double(&nxt) // nxt <- 2 * qᵢ * hᵢ - nxt.Add(&nxt, &one) // nxt <- 1 + 2 * qᵢ * hᵢ - sum.Add(&q[i], &h[i]) // sum <- qᵢ + hᵢ TODO: Why not subtract one by one from nxt? More parallel? - - if i == 0 { - res.Sub(&nxt, &sum) // nxt <- 1 + 2 * qᵢ * hᵢ - qᵢ - hᵢ - } else { - nxt.Sub(&nxt, &sum) // nxt <- 1 + 2 * qᵢ * hᵢ - qᵢ - hᵢ - res.Mul(&res, &nxt) // res <- res * nxt - } - } - return res -} - -// Eq sets m to the representation of the polynomial Eq(q₁, ..., qₙ, *, ..., *) × m[0] -func (m *MultiLin) Eq(q []babybear.Element) { - n := len(q) - - if len(*m) != 1<= 0; i-- { - res.Mul(&res, v) - res.Add(&res, &(*p)[i]) - } - - return res -} - -// Clone returns a copy of the polynomial -func (p *Polynomial) Clone() Polynomial { - _p := make(Polynomial, len(*p)) - copy(_p, *p) - return _p -} - -// Set to another polynomial -func (p *Polynomial) Set(p1 Polynomial) { - if len(*p) != len(p1) { - *p = p1.Clone() - return - } - - for i := range len(p1) { - (*p)[i].Set(&p1[i]) - } -} - -// AddConstantInPlace adds a constant to the polynomial, modifying p -func (p *Polynomial) AddConstantInPlace(c *babybear.Element) { - for i := range len(*p) { - (*p)[i].Add(&(*p)[i], c) - } -} - -// SubConstantInPlace subs a constant to the polynomial, modifying p -func (p *Polynomial) SubConstantInPlace(c *babybear.Element) { - for i := range len(*p) { - (*p)[i].Sub(&(*p)[i], c) - } -} - -// ScaleInPlace multiplies p by v, modifying p -func (p *Polynomial) ScaleInPlace(c *babybear.Element) { - for i := range len(*p) { - (*p)[i].Mul(&(*p)[i], c) - } -} - -// Scale multiplies p0 by v, storing the result in p -func (p *Polynomial) Scale(c *babybear.Element, p0 Polynomial) { - if len(*p) != len(p0) { - *p = make(Polynomial, len(p0)) - } - for i := range len(p0) { - (*p)[i].Mul(c, &p0[i]) - } -} - -// Add adds p1 to p2 -// This function allocates a new slice unless p == p1 or p == p2 -func (p *Polynomial) Add(p1, p2 Polynomial) *Polynomial { - - bigger := p1 - smaller := p2 - if len(bigger) < len(smaller) { - bigger, smaller = smaller, bigger - } - - if len(*p) == len(bigger) && (&(*p)[0] == &bigger[0]) { - for i := range len(smaller) { - (*p)[i].Add(&(*p)[i], &smaller[i]) - } - return p - } - - if len(*p) == len(smaller) && (&(*p)[0] == &smaller[0]) { - for i := range len(smaller) { - (*p)[i].Add(&(*p)[i], &bigger[i]) - } - *p = append(*p, bigger[len(smaller):]...) - return p - } - - res := make(Polynomial, len(bigger)) - copy(res, bigger) - for i := range len(smaller) { - res[i].Add(&res[i], &smaller[i]) - } - *p = res - return p -} - -// Sub subtracts p2 from p1 -// TODO make interface more consistent with Add -func (p *Polynomial) Sub(p1, p2 Polynomial) *Polynomial { - if len(p1) != len(p2) || len(p2) != len(*p) { - return nil - } - for i := range len(*p) { - (*p)[i].Sub(&p1[i], &p2[i]) - } - return p -} - -// Equal checks equality between two polynomials -func (p *Polynomial) Equal(p1 Polynomial) bool { - if (*p == nil) != (p1 == nil) { - return false - } - - if len(*p) != len(p1) { - return false - } - - for i := range p1 { - if !(*p)[i].Equal(&p1[i]) { - return false - } - } - - return true -} - -func (p Polynomial) SetZero() { - for i := range len(p) { - p[i].SetZero() - } -} - -func (p Polynomial) Text(base int) string { - - var builder strings.Builder - - first := true - for d := len(p) - 1; d >= 0; d-- { - if p[d].IsZero() { - continue - } - - pD := p[d] - pDText := pD.Text(base) - - initialLen := builder.Len() - - if pDText[0] == '-' { - pDText = pDText[1:] - if first { - builder.WriteString("-") - } else { - builder.WriteString(" - ") - } - } else if !first { - builder.WriteString(" + ") - } - - first = false - - if !pD.IsOne() || d == 0 { - builder.WriteString(pDText) - } - - if builder.Len()-initialLen > 10 { - builder.WriteString("×") - } - - if d != 0 { - builder.WriteString("X") - } - if d > 1 { - builder.WriteString( - utils.ToSuperscript(strconv.Itoa(d)), - ) - } - - } - - if first { - return "0" - } - - return builder.String() -} - -// InterpolateOnRange maps vector v to polynomial f -// such that f(i) = v[i] for 0 ≤ i < len(v). -// len(f) = len(v) and deg(f) ≤ len(v) - 1 -func InterpolateOnRange(v []babybear.Element) Polynomial { - nEvals := uint8(len(v)) - if int(nEvals) != len(v) { - panic("interpolation method too inefficient for nEvals > 255") - } - lagrange := getLagrangeBasis(nEvals) - - var res Polynomial - res.Scale(&v[0], lagrange[0]) - - temp := make(Polynomial, nEvals) - - for i := uint8(1); i < nEvals; i++ { - temp.Scale(&v[i], lagrange[i]) - res.Add(res, temp) - } - - return res -} - -// lagrange bases used by InterpolateOnRange -var lagrangeBasis sync.Map - -func getLagrangeBasis(domainSize uint8) []Polynomial { - if res, ok := lagrangeBasis.Load(domainSize); ok { - return res.([]Polynomial) - } - - // not found. compute - var res []Polynomial - if domainSize >= 2 { - res = computeLagrangeBasis(domainSize) - } else if domainSize == 1 { - res = []Polynomial{make(Polynomial, 1)} - res[0][0].SetOne() - } - lagrangeBasis.Store(domainSize, res) - - return res -} - -// computeLagrangeBasis precomputes in explicit coefficient form for each 0 ≤ l < domainSize the polynomial -// pₗ := X (X-1) ... (X-l-1) (X-l+1) ... (X - domainSize + 1) / ( l (l-1) ... 2 (-1) ... (l - domainSize +1) ) -// Note that pₗ(l) = 1 and pₗ(n) = 0 if 0 ≤ l < domainSize, n ≠ l -func computeLagrangeBasis(domainSize uint8) []Polynomial { - - constTerms := make([]babybear.Element, domainSize) - for i := range domainSize { - constTerms[i].SetInt64(-int64(i)) - } - - res := make([]Polynomial, domainSize) - multScratch := make(Polynomial, domainSize-1) - - // compute pₗ - for l := range domainSize { - - // TODO @Tabaie Optimize this with some trees? O(log(domainSize)) polynomial mults instead of O(domainSize)? Then again it would be fewer big poly mults vs many small poly mults - d := uint8(0) //d is the current degree of res - for i := range domainSize { - if i == l { - continue - } - if d == 0 { - res[l] = make(Polynomial, domainSize) - res[l][domainSize-2] = constTerms[i] - res[l][domainSize-1].SetOne() - } else { - current := res[l][domainSize-d-2:] - timesConst := multScratch[domainSize-d-2:] - - timesConst.Scale(&constTerms[i], current[1:]) //TODO: Directly double and add since constTerms are tiny? (even less than 4 bits) - nonLeading := current[0 : d+1] - - nonLeading.Add(nonLeading, timesConst) - - } - d++ - } - - } - - // We have pₗ(i≠l)=0. Now scale so that pₗ(l)=1 - // Replace the constTerms with norms - for l := range domainSize { - constTerms[l].Neg(&constTerms[l]) - constTerms[l] = res[l].Eval(&constTerms[l]) - } - constTerms = babybear.BatchInvert(constTerms) - for l := range domainSize { - res[l].ScaleInPlace(&constTerms[l]) - } - - return res -} diff --git a/field/babybear/polynomial/polynomial_test.go b/field/babybear/polynomial/polynomial_test.go deleted file mode 100644 index efcac7a903..0000000000 --- a/field/babybear/polynomial/polynomial_test.go +++ /dev/null @@ -1,255 +0,0 @@ -// Copyright 2020-2026 Consensys Software Inc. -// Licensed under the Apache License, Version 2.0. See the LICENSE file for details. - -// Code generated by consensys/gnark-crypto DO NOT EDIT - -package polynomial - -import ( - "math/big" - "testing" - - "github.com/consensys/gnark-crypto/field/babybear" - "github.com/leanovate/gopter" - "github.com/leanovate/gopter/gen" - "github.com/leanovate/gopter/prop" - "github.com/stretchr/testify/assert" -) - -func TestPolynomialEval(t *testing.T) { - - // build polynomial - f := make(Polynomial, 20) - for i := range 20 { - f[i].SetOne() - } - - // random value - var point babybear.Element - point.MustSetRandom() - - // compute manually f(val) - var expectedEval, one, den babybear.Element - var expo big.Int - one.SetOne() - expo.SetUint64(20) - expectedEval.Exp(point, &expo). - Sub(&expectedEval, &one) - den.Sub(&point, &one) - expectedEval.Div(&expectedEval, &den) - - // compute purported evaluation - purportedEval := f.Eval(&point) - - // check - if !purportedEval.Equal(&expectedEval) { - t.Fatal("polynomial evaluation failed") - } -} - -func TestPolynomialAddConstantInPlace(t *testing.T) { - - // build polynomial - f := make(Polynomial, 20) - for i := range 20 { - f[i].SetOne() - } - - // constant to add - var c babybear.Element - c.MustSetRandom() - - // add constant - f.AddConstantInPlace(&c) - - // check - var expectedCoeffs, one babybear.Element - one.SetOne() - expectedCoeffs.Add(&one, &c) - for i := range 20 { - if !f[i].Equal(&expectedCoeffs) { - t.Fatal("AddConstantInPlace failed") - } - } -} - -func TestPolynomialSubConstantInPlace(t *testing.T) { - - // build polynomial - f := make(Polynomial, 20) - for i := range 20 { - f[i].SetOne() - } - - // constant to sub - var c babybear.Element - c.MustSetRandom() - - // sub constant - f.SubConstantInPlace(&c) - - // check - var expectedCoeffs, one babybear.Element - one.SetOne() - expectedCoeffs.Sub(&one, &c) - for i := range 20 { - if !f[i].Equal(&expectedCoeffs) { - t.Fatal("SubConstantInPlace failed") - } - } -} - -func TestPolynomialScaleInPlace(t *testing.T) { - - // build polynomial - f := make(Polynomial, 20) - for i := range 20 { - f[i].SetOne() - } - - // constant to scale by - var c babybear.Element - c.MustSetRandom() - - // scale by constant - f.ScaleInPlace(&c) - - // check - for i := range 20 { - if !f[i].Equal(&c) { - t.Fatal("ScaleInPlace failed") - } - } - -} - -func TestPolynomialAdd(t *testing.T) { - - // build unbalanced polynomials - f1 := make(Polynomial, 20) - f1Backup := make(Polynomial, 20) - for i := range 20 { - f1[i].SetOne() - f1Backup[i].SetOne() - } - f2 := make(Polynomial, 10) - f2Backup := make(Polynomial, 10) - for i := range 10 { - f2[i].SetOne() - f2Backup[i].SetOne() - } - - // expected result - var one, two babybear.Element - one.SetOne() - two.Double(&one) - expectedSum := make(Polynomial, 20) - for i := range 10 { - expectedSum[i].Set(&two) - } - for i := 10; i < 20; i++ { - expectedSum[i].Set(&one) - } - - // caller is empty - var g Polynomial - g.Add(f1, f2) - if !g.Equal(expectedSum) { - t.Fatal("add polynomials fails") - } - if !f1.Equal(f1Backup) { - t.Fatal("side effect, f1 should not have been modified") - } - if !f2.Equal(f2Backup) { - t.Fatal("side effect, f2 should not have been modified") - } - - // all operands are distinct - _f1 := f1.Clone() - _f1.Add(f1, f2) - if !_f1.Equal(expectedSum) { - t.Fatal("add polynomials fails") - } - if !f1.Equal(f1Backup) { - t.Fatal("side effect, f1 should not have been modified") - } - if !f2.Equal(f2Backup) { - t.Fatal("side effect, f2 should not have been modified") - } - - // first operand = caller - _f1 = f1.Clone() - _f2 := f2.Clone() - _f1.Add(_f1, _f2) - if !_f1.Equal(expectedSum) { - t.Fatal("add polynomials fails") - } - if !_f2.Equal(f2Backup) { - t.Fatal("side effect, _f2 should not have been modified") - } - - // second operand = caller - _f1 = f1.Clone() - _f2 = f2.Clone() - _f1.Add(_f2, _f1) - if !_f1.Equal(expectedSum) { - t.Fatal("add polynomials fails") - } - if !_f2.Equal(f2Backup) { - t.Fatal("side effect, _f2 should not have been modified") - } -} - -func TestPolynomialText(t *testing.T) { - var one, negTwo babybear.Element - one.SetOne() - negTwo.SetInt64(-2) - - p := Polynomial{one, negTwo, one} - - assert.Equal(t, "X² - 2X + 1", p.Text(10)) -} - -func TestPrecomputeLagrange(t *testing.T) { - - testForDomainSize := func(domainSize uint8) bool { - polys := computeLagrangeBasis(domainSize) - - for l := range domainSize { - for i := range domainSize { - var I babybear.Element - I.SetUint64(uint64(i)) - y := polys[l].Eval(&I) - - if i == l && !y.IsOne() || i != l && !y.IsZero() { - t.Errorf("domainSize = %d: p_%d(%d) = %s", domainSize, l, i, y.Text(10)) - return false - } - } - } - return true - } - - t.Parallel() - parameters := gopter.DefaultTestParameters() - - const maxLagrangeDomainSize = 12 - - parameters.MinSuccessfulTests = maxLagrangeDomainSize - - properties := gopter.NewProperties(parameters) - - properties.Property("l'th lagrange polynomials must evaluate to 1 on l and 0 on other values in the domain", prop.ForAll( - testForDomainSize, - gen.UInt8Range(2, maxLagrangeDomainSize), - )) - - properties.TestingRun(t, gopter.ConsoleReporter(false)) -} - -func TestLagrangeCache(t *testing.T) { - for _, i := range []int{5, 2, 8, 4, 6, 3, 0} { - b := getLagrangeBasis(uint8(i)) - assert.Equal(t, b, getLagrangeBasis(uint8(i))) // second call must yield the same result - } -} diff --git a/field/babybear/polynomial/pool.go b/field/babybear/polynomial/pool.go deleted file mode 100644 index 6a03e14325..0000000000 --- a/field/babybear/polynomial/pool.go +++ /dev/null @@ -1,191 +0,0 @@ -// Copyright 2020-2026 Consensys Software Inc. -// Licensed under the Apache License, Version 2.0. See the LICENSE file for details. - -// Code generated by consensys/gnark-crypto DO NOT EDIT - -package polynomial - -import ( - "encoding/json" - "fmt" - "runtime" - "sort" - "sync" - "unsafe" - - "github.com/consensys/gnark-crypto/field/babybear" -) - -// Memory management for polynomials -// WARNING: This is not thread safe TODO: Make sure that is not a problem -// TODO: There is a lot of "unsafe" memory management here and needs to be vetted thoroughly - -type sizedPool struct { - maxN int - pool sync.Pool - stats poolStats -} - -type inUseData struct { - allocatedFor []uintptr - pool *sizedPool -} - -type Pool struct { - //lock sync.Mutex - inUse sync.Map - subPools []sizedPool -} - -func (p *sizedPool) get(n int) *babybear.Element { - p.stats.make(n) - return p.pool.Get().(*babybear.Element) -} - -func (p *sizedPool) put(ptr *babybear.Element) { - p.stats.dump() - p.pool.Put(ptr) -} - -func NewPool(maxN ...int) (pool Pool) { - - sort.Ints(maxN) - pool = Pool{ - subPools: make([]sizedPool, len(maxN)), - } - - for i := range pool.subPools { - subPool := &pool.subPools[i] - subPool.maxN = maxN[i] - subPool.pool = sync.Pool{ - New: func() any { - subPool.stats.Allocated++ - return getDataPointer(make([]babybear.Element, 0, subPool.maxN)) - }, - } - } - return -} - -func (p *Pool) findCorrespondingPool(n int) *sizedPool { - poolI := 0 - for poolI < len(p.subPools) && n > p.subPools[poolI].maxN { - poolI++ - } - return &p.subPools[poolI] // out of bounds error here would mean that n is too large -} - -func (p *Pool) Make(n int) []babybear.Element { - pool := p.findCorrespondingPool(n) - ptr := pool.get(n) - p.addInUse(ptr, pool) - return unsafe.Slice(ptr, n) -} - -// Dump dumps a set of polynomials into the pool -func (p *Pool) Dump(slices ...[]babybear.Element) { - for _, slice := range slices { - ptr := getDataPointer(slice) - if metadata, ok := p.inUse.Load(ptr); ok { - p.inUse.Delete(ptr) - metadata.(inUseData).pool.put(ptr) - } else { - panic("attempting to dump a slice not created by the pool") - } - } -} - -func (p *Pool) addInUse(ptr *babybear.Element, pool *sizedPool) { - pcs := make([]uintptr, 2) - n := runtime.Callers(3, pcs) - - if prevPcs, ok := p.inUse.Load(ptr); ok { // TODO: remove if unnecessary for security - panic(fmt.Errorf("re-allocated non-dumped slice, previously allocated at %v", runtime.CallersFrames(prevPcs.(inUseData).allocatedFor))) - } - p.inUse.Store(ptr, inUseData{ - allocatedFor: pcs[:n], - pool: pool, - }) -} - -func printFrame(frame runtime.Frame) { - fmt.Printf("\t%s line %d, function %s\n", frame.File, frame.Line, frame.Function) -} - -func (p *Pool) printInUse() { - fmt.Println("slices never dumped allocated at:") - p.inUse.Range(func(_, pcs any) bool { - fmt.Println("-------------------------") - - var frame runtime.Frame - frames := runtime.CallersFrames(pcs.(inUseData).allocatedFor) - more := true - for more { - frame, more = frames.Next() - printFrame(frame) - } - return true - }) -} - -type poolStats struct { - Used int - Allocated int - ReuseRate float64 - InUse int - GreatestNUsed int - SmallestNUsed int -} - -type poolsStats struct { - SubPools []poolStats - InUse int -} - -func (s *poolStats) make(n int) { - s.Used++ - s.InUse++ - if n > s.GreatestNUsed { - s.GreatestNUsed = n - } - if s.SmallestNUsed == 0 || s.SmallestNUsed > n { - s.SmallestNUsed = n - } -} - -func (s *poolStats) dump() { - s.InUse-- -} - -func (s *poolStats) finalize() { - s.ReuseRate = float64(s.Used) / float64(s.Allocated) -} - -func getDataPointer(slice []babybear.Element) *babybear.Element { - return (*babybear.Element)(unsafe.SliceData(slice)) -} - -func (p *Pool) PrintPoolStats() { - InUse := 0 - subStats := make([]poolStats, len(p.subPools)) - for i := range p.subPools { - subPool := &p.subPools[i] - subPool.stats.finalize() - subStats[i] = subPool.stats - InUse += subPool.stats.InUse - } - - stats := poolsStats{ - SubPools: subStats, - InUse: InUse, - } - serialized, _ := json.MarshalIndent(stats, "", " ") - fmt.Println(string(serialized)) - p.printInUse() -} - -func (p *Pool) Clone(slice []babybear.Element) []babybear.Element { - res := p.Make(len(slice)) - copy(res, slice) - return res -} diff --git a/field/goldilocks/polynomial/doc.go b/field/goldilocks/polynomial/doc.go deleted file mode 100644 index aa346f3ea3..0000000000 --- a/field/goldilocks/polynomial/doc.go +++ /dev/null @@ -1,7 +0,0 @@ -// Copyright 2020-2026 Consensys Software Inc. -// Licensed under the Apache License, Version 2.0. See the LICENSE file for details. - -// Code generated by consensys/gnark-crypto DO NOT EDIT - -// Package polynomial provides polynomial methods and commitment schemes. -package polynomial diff --git a/field/goldilocks/polynomial/multilin.go b/field/goldilocks/polynomial/multilin.go deleted file mode 100644 index bfcc0f96f6..0000000000 --- a/field/goldilocks/polynomial/multilin.go +++ /dev/null @@ -1,179 +0,0 @@ -// Copyright 2020-2026 Consensys Software Inc. -// Licensed under the Apache License, Version 2.0. See the LICENSE file for details. - -// Code generated by consensys/gnark-crypto DO NOT EDIT - -package polynomial - -import ( - "math/bits" - - "github.com/consensys/gnark-crypto/field/goldilocks" - "github.com/consensys/gnark-crypto/utils" -) - -// MultiLin tracks the values of a (dense i.e. not sparse) multilinear polynomial -// The variables are X₁ through Xₙ where n = log(len(.)) -// .[∑ᵢ 2ⁱ⁻¹ bₙ₋ᵢ] = the polynomial evaluated at (b₁, b₂, ..., bₙ) -// It is understood that any hypercube evaluation can be extrapolated to a multilinear polynomial -type MultiLin []goldilocks.Element - -// Fold is partial evaluation function k[X₁, X₂, ..., Xₙ] → k[X₂, ..., Xₙ] by setting X₁=r -func (m *MultiLin) Fold(r goldilocks.Element) { - mid := len(*m) / 2 - - bottom, top := (*m)[:mid], (*m)[mid:] - - var t goldilocks.Element // no need to update the top part - - // updating bookkeeping table - // knowing that the polynomial f ∈ (k[X₂, ..., Xₙ])[X₁] is linear, we would get f(r) = f(0) + r(f(1) - f(0)) - // the following loop computes the evaluations of f(r) accordingly: - // f(r, b₂, ..., bₙ) = f(0, b₂, ..., bₙ) + r(f(1, b₂, ..., bₙ) - f(0, b₂, ..., bₙ)) - for i := range mid { - // table[i] ← table[i] + r (table[i + mid] - table[i]) - t.Sub(&top[i], &bottom[i]) - t.Mul(&t, &r) - bottom[i].Add(&bottom[i], &t) - } - - *m = (*m)[:mid] -} - -func (m *MultiLin) FoldParallel(r goldilocks.Element) utils.Task { - mid := len(*m) / 2 - bottom, top := (*m)[:mid], (*m)[mid:] - - *m = bottom - - return func(start, end int) { - var t goldilocks.Element // no need to update the top part - for i := start; i < end; i++ { - // table[i] ← table[i] + r (table[i + mid] - table[i]) - t.Sub(&top[i], &bottom[i]) - t.Mul(&t, &r) - bottom[i].Add(&bottom[i], &t) - } - } -} - -func (m MultiLin) Sum() goldilocks.Element { - s := m[0] - for i := 1; i < len(m); i++ { - s.Add(&s, &m[i]) - } - return s -} - -func _clone(m MultiLin, p *Pool) MultiLin { - if p == nil { - return m.Clone() - } else { - return p.Clone(m) - } -} - -func _dump(m MultiLin, p *Pool) { - if p != nil { - p.Dump(m) - } -} - -// Evaluate extrapolate the value of the multilinear polynomial corresponding to m -// on the given coordinates -func (m MultiLin) Evaluate(coordinates []goldilocks.Element, p *Pool) goldilocks.Element { - // Folding is a mutating operation - bkCopy := _clone(m, p) - - // Evaluate step by step through repeated folding (i.e. evaluation at the first remaining variable) - for _, r := range coordinates { - bkCopy.Fold(r) - } - - result := bkCopy[0] - - _dump(bkCopy, p) - return result -} - -// Clone creates a deep copy of a bookkeeping table. -// Both multilinear interpolation and sumcheck require folding an underlying -// array, but folding changes the array. To do both one requires a deep copy -// of the bookkeeping table. -func (m MultiLin) Clone() MultiLin { - res := make(MultiLin, len(m)) - copy(res, m) - return res -} - -// Add two bookKeepingTables -func (m *MultiLin) Add(left, right MultiLin) { - size := len(left) - // Check that left and right have the same size - if len(right) != size || len(*m) != size { - panic("left, right and destination must have the right size") - } - - // Add elementwise - for i := range size { - (*m)[i].Add(&left[i], &right[i]) - } -} - -// EvalEq computes Eq(q₁, ... , qₙ, h₁, ... , hₙ) = Π₁ⁿ Eq(qᵢ, hᵢ) -// where Eq(x,y) = xy + (1-x)(1-y) = 1 - x - y + xy + xy interpolates -// -// _________________ -// | | | -// | 0 | 1 | -// |_______|_______| -// y | | | -// | 1 | 0 | -// |_______|_______| -// -// x -// -// In other words the polynomial evaluated here is the multilinear extrapolation of -// one that evaluates to q' == h' for vectors q', h' of binary values -func EvalEq(q, h []goldilocks.Element) goldilocks.Element { - var res, nxt, one, sum goldilocks.Element - one.SetOne() - for i := range len(q) { - nxt.Mul(&q[i], &h[i]) // nxt <- qᵢ * hᵢ - nxt.Double(&nxt) // nxt <- 2 * qᵢ * hᵢ - nxt.Add(&nxt, &one) // nxt <- 1 + 2 * qᵢ * hᵢ - sum.Add(&q[i], &h[i]) // sum <- qᵢ + hᵢ TODO: Why not subtract one by one from nxt? More parallel? - - if i == 0 { - res.Sub(&nxt, &sum) // nxt <- 1 + 2 * qᵢ * hᵢ - qᵢ - hᵢ - } else { - nxt.Sub(&nxt, &sum) // nxt <- 1 + 2 * qᵢ * hᵢ - qᵢ - hᵢ - res.Mul(&res, &nxt) // res <- res * nxt - } - } - return res -} - -// Eq sets m to the representation of the polynomial Eq(q₁, ..., qₙ, *, ..., *) × m[0] -func (m *MultiLin) Eq(q []goldilocks.Element) { - n := len(q) - - if len(*m) != 1<= 0; i-- { - res.Mul(&res, v) - res.Add(&res, &(*p)[i]) - } - - return res -} - -// Clone returns a copy of the polynomial -func (p *Polynomial) Clone() Polynomial { - _p := make(Polynomial, len(*p)) - copy(_p, *p) - return _p -} - -// Set to another polynomial -func (p *Polynomial) Set(p1 Polynomial) { - if len(*p) != len(p1) { - *p = p1.Clone() - return - } - - for i := range len(p1) { - (*p)[i].Set(&p1[i]) - } -} - -// AddConstantInPlace adds a constant to the polynomial, modifying p -func (p *Polynomial) AddConstantInPlace(c *goldilocks.Element) { - for i := range len(*p) { - (*p)[i].Add(&(*p)[i], c) - } -} - -// SubConstantInPlace subs a constant to the polynomial, modifying p -func (p *Polynomial) SubConstantInPlace(c *goldilocks.Element) { - for i := range len(*p) { - (*p)[i].Sub(&(*p)[i], c) - } -} - -// ScaleInPlace multiplies p by v, modifying p -func (p *Polynomial) ScaleInPlace(c *goldilocks.Element) { - for i := range len(*p) { - (*p)[i].Mul(&(*p)[i], c) - } -} - -// Scale multiplies p0 by v, storing the result in p -func (p *Polynomial) Scale(c *goldilocks.Element, p0 Polynomial) { - if len(*p) != len(p0) { - *p = make(Polynomial, len(p0)) - } - for i := range len(p0) { - (*p)[i].Mul(c, &p0[i]) - } -} - -// Add adds p1 to p2 -// This function allocates a new slice unless p == p1 or p == p2 -func (p *Polynomial) Add(p1, p2 Polynomial) *Polynomial { - - bigger := p1 - smaller := p2 - if len(bigger) < len(smaller) { - bigger, smaller = smaller, bigger - } - - if len(*p) == len(bigger) && (&(*p)[0] == &bigger[0]) { - for i := range len(smaller) { - (*p)[i].Add(&(*p)[i], &smaller[i]) - } - return p - } - - if len(*p) == len(smaller) && (&(*p)[0] == &smaller[0]) { - for i := range len(smaller) { - (*p)[i].Add(&(*p)[i], &bigger[i]) - } - *p = append(*p, bigger[len(smaller):]...) - return p - } - - res := make(Polynomial, len(bigger)) - copy(res, bigger) - for i := range len(smaller) { - res[i].Add(&res[i], &smaller[i]) - } - *p = res - return p -} - -// Sub subtracts p2 from p1 -// TODO make interface more consistent with Add -func (p *Polynomial) Sub(p1, p2 Polynomial) *Polynomial { - if len(p1) != len(p2) || len(p2) != len(*p) { - return nil - } - for i := range len(*p) { - (*p)[i].Sub(&p1[i], &p2[i]) - } - return p -} - -// Equal checks equality between two polynomials -func (p *Polynomial) Equal(p1 Polynomial) bool { - if (*p == nil) != (p1 == nil) { - return false - } - - if len(*p) != len(p1) { - return false - } - - for i := range p1 { - if !(*p)[i].Equal(&p1[i]) { - return false - } - } - - return true -} - -func (p Polynomial) SetZero() { - for i := range len(p) { - p[i].SetZero() - } -} - -func (p Polynomial) Text(base int) string { - - var builder strings.Builder - - first := true - for d := len(p) - 1; d >= 0; d-- { - if p[d].IsZero() { - continue - } - - pD := p[d] - pDText := pD.Text(base) - - initialLen := builder.Len() - - if pDText[0] == '-' { - pDText = pDText[1:] - if first { - builder.WriteString("-") - } else { - builder.WriteString(" - ") - } - } else if !first { - builder.WriteString(" + ") - } - - first = false - - if !pD.IsOne() || d == 0 { - builder.WriteString(pDText) - } - - if builder.Len()-initialLen > 10 { - builder.WriteString("×") - } - - if d != 0 { - builder.WriteString("X") - } - if d > 1 { - builder.WriteString( - utils.ToSuperscript(strconv.Itoa(d)), - ) - } - - } - - if first { - return "0" - } - - return builder.String() -} - -// InterpolateOnRange maps vector v to polynomial f -// such that f(i) = v[i] for 0 ≤ i < len(v). -// len(f) = len(v) and deg(f) ≤ len(v) - 1 -func InterpolateOnRange(v []goldilocks.Element) Polynomial { - nEvals := uint8(len(v)) - if int(nEvals) != len(v) { - panic("interpolation method too inefficient for nEvals > 255") - } - lagrange := getLagrangeBasis(nEvals) - - var res Polynomial - res.Scale(&v[0], lagrange[0]) - - temp := make(Polynomial, nEvals) - - for i := uint8(1); i < nEvals; i++ { - temp.Scale(&v[i], lagrange[i]) - res.Add(res, temp) - } - - return res -} - -// lagrange bases used by InterpolateOnRange -var lagrangeBasis sync.Map - -func getLagrangeBasis(domainSize uint8) []Polynomial { - if res, ok := lagrangeBasis.Load(domainSize); ok { - return res.([]Polynomial) - } - - // not found. compute - var res []Polynomial - if domainSize >= 2 { - res = computeLagrangeBasis(domainSize) - } else if domainSize == 1 { - res = []Polynomial{make(Polynomial, 1)} - res[0][0].SetOne() - } - lagrangeBasis.Store(domainSize, res) - - return res -} - -// computeLagrangeBasis precomputes in explicit coefficient form for each 0 ≤ l < domainSize the polynomial -// pₗ := X (X-1) ... (X-l-1) (X-l+1) ... (X - domainSize + 1) / ( l (l-1) ... 2 (-1) ... (l - domainSize +1) ) -// Note that pₗ(l) = 1 and pₗ(n) = 0 if 0 ≤ l < domainSize, n ≠ l -func computeLagrangeBasis(domainSize uint8) []Polynomial { - - constTerms := make([]goldilocks.Element, domainSize) - for i := range domainSize { - constTerms[i].SetInt64(-int64(i)) - } - - res := make([]Polynomial, domainSize) - multScratch := make(Polynomial, domainSize-1) - - // compute pₗ - for l := range domainSize { - - // TODO @Tabaie Optimize this with some trees? O(log(domainSize)) polynomial mults instead of O(domainSize)? Then again it would be fewer big poly mults vs many small poly mults - d := uint8(0) //d is the current degree of res - for i := range domainSize { - if i == l { - continue - } - if d == 0 { - res[l] = make(Polynomial, domainSize) - res[l][domainSize-2] = constTerms[i] - res[l][domainSize-1].SetOne() - } else { - current := res[l][domainSize-d-2:] - timesConst := multScratch[domainSize-d-2:] - - timesConst.Scale(&constTerms[i], current[1:]) //TODO: Directly double and add since constTerms are tiny? (even less than 4 bits) - nonLeading := current[0 : d+1] - - nonLeading.Add(nonLeading, timesConst) - - } - d++ - } - - } - - // We have pₗ(i≠l)=0. Now scale so that pₗ(l)=1 - // Replace the constTerms with norms - for l := range domainSize { - constTerms[l].Neg(&constTerms[l]) - constTerms[l] = res[l].Eval(&constTerms[l]) - } - constTerms = goldilocks.BatchInvert(constTerms) - for l := range domainSize { - res[l].ScaleInPlace(&constTerms[l]) - } - - return res -} diff --git a/field/goldilocks/polynomial/polynomial_test.go b/field/goldilocks/polynomial/polynomial_test.go deleted file mode 100644 index 3a36ae3a53..0000000000 --- a/field/goldilocks/polynomial/polynomial_test.go +++ /dev/null @@ -1,255 +0,0 @@ -// Copyright 2020-2026 Consensys Software Inc. -// Licensed under the Apache License, Version 2.0. See the LICENSE file for details. - -// Code generated by consensys/gnark-crypto DO NOT EDIT - -package polynomial - -import ( - "math/big" - "testing" - - "github.com/consensys/gnark-crypto/field/goldilocks" - "github.com/leanovate/gopter" - "github.com/leanovate/gopter/gen" - "github.com/leanovate/gopter/prop" - "github.com/stretchr/testify/assert" -) - -func TestPolynomialEval(t *testing.T) { - - // build polynomial - f := make(Polynomial, 20) - for i := range 20 { - f[i].SetOne() - } - - // random value - var point goldilocks.Element - point.MustSetRandom() - - // compute manually f(val) - var expectedEval, one, den goldilocks.Element - var expo big.Int - one.SetOne() - expo.SetUint64(20) - expectedEval.Exp(point, &expo). - Sub(&expectedEval, &one) - den.Sub(&point, &one) - expectedEval.Div(&expectedEval, &den) - - // compute purported evaluation - purportedEval := f.Eval(&point) - - // check - if !purportedEval.Equal(&expectedEval) { - t.Fatal("polynomial evaluation failed") - } -} - -func TestPolynomialAddConstantInPlace(t *testing.T) { - - // build polynomial - f := make(Polynomial, 20) - for i := range 20 { - f[i].SetOne() - } - - // constant to add - var c goldilocks.Element - c.MustSetRandom() - - // add constant - f.AddConstantInPlace(&c) - - // check - var expectedCoeffs, one goldilocks.Element - one.SetOne() - expectedCoeffs.Add(&one, &c) - for i := range 20 { - if !f[i].Equal(&expectedCoeffs) { - t.Fatal("AddConstantInPlace failed") - } - } -} - -func TestPolynomialSubConstantInPlace(t *testing.T) { - - // build polynomial - f := make(Polynomial, 20) - for i := range 20 { - f[i].SetOne() - } - - // constant to sub - var c goldilocks.Element - c.MustSetRandom() - - // sub constant - f.SubConstantInPlace(&c) - - // check - var expectedCoeffs, one goldilocks.Element - one.SetOne() - expectedCoeffs.Sub(&one, &c) - for i := range 20 { - if !f[i].Equal(&expectedCoeffs) { - t.Fatal("SubConstantInPlace failed") - } - } -} - -func TestPolynomialScaleInPlace(t *testing.T) { - - // build polynomial - f := make(Polynomial, 20) - for i := range 20 { - f[i].SetOne() - } - - // constant to scale by - var c goldilocks.Element - c.MustSetRandom() - - // scale by constant - f.ScaleInPlace(&c) - - // check - for i := range 20 { - if !f[i].Equal(&c) { - t.Fatal("ScaleInPlace failed") - } - } - -} - -func TestPolynomialAdd(t *testing.T) { - - // build unbalanced polynomials - f1 := make(Polynomial, 20) - f1Backup := make(Polynomial, 20) - for i := range 20 { - f1[i].SetOne() - f1Backup[i].SetOne() - } - f2 := make(Polynomial, 10) - f2Backup := make(Polynomial, 10) - for i := range 10 { - f2[i].SetOne() - f2Backup[i].SetOne() - } - - // expected result - var one, two goldilocks.Element - one.SetOne() - two.Double(&one) - expectedSum := make(Polynomial, 20) - for i := range 10 { - expectedSum[i].Set(&two) - } - for i := 10; i < 20; i++ { - expectedSum[i].Set(&one) - } - - // caller is empty - var g Polynomial - g.Add(f1, f2) - if !g.Equal(expectedSum) { - t.Fatal("add polynomials fails") - } - if !f1.Equal(f1Backup) { - t.Fatal("side effect, f1 should not have been modified") - } - if !f2.Equal(f2Backup) { - t.Fatal("side effect, f2 should not have been modified") - } - - // all operands are distinct - _f1 := f1.Clone() - _f1.Add(f1, f2) - if !_f1.Equal(expectedSum) { - t.Fatal("add polynomials fails") - } - if !f1.Equal(f1Backup) { - t.Fatal("side effect, f1 should not have been modified") - } - if !f2.Equal(f2Backup) { - t.Fatal("side effect, f2 should not have been modified") - } - - // first operand = caller - _f1 = f1.Clone() - _f2 := f2.Clone() - _f1.Add(_f1, _f2) - if !_f1.Equal(expectedSum) { - t.Fatal("add polynomials fails") - } - if !_f2.Equal(f2Backup) { - t.Fatal("side effect, _f2 should not have been modified") - } - - // second operand = caller - _f1 = f1.Clone() - _f2 = f2.Clone() - _f1.Add(_f2, _f1) - if !_f1.Equal(expectedSum) { - t.Fatal("add polynomials fails") - } - if !_f2.Equal(f2Backup) { - t.Fatal("side effect, _f2 should not have been modified") - } -} - -func TestPolynomialText(t *testing.T) { - var one, negTwo goldilocks.Element - one.SetOne() - negTwo.SetInt64(-2) - - p := Polynomial{one, negTwo, one} - - assert.Equal(t, "X² - 2X + 1", p.Text(10)) -} - -func TestPrecomputeLagrange(t *testing.T) { - - testForDomainSize := func(domainSize uint8) bool { - polys := computeLagrangeBasis(domainSize) - - for l := range domainSize { - for i := range domainSize { - var I goldilocks.Element - I.SetUint64(uint64(i)) - y := polys[l].Eval(&I) - - if i == l && !y.IsOne() || i != l && !y.IsZero() { - t.Errorf("domainSize = %d: p_%d(%d) = %s", domainSize, l, i, y.Text(10)) - return false - } - } - } - return true - } - - t.Parallel() - parameters := gopter.DefaultTestParameters() - - const maxLagrangeDomainSize = 12 - - parameters.MinSuccessfulTests = maxLagrangeDomainSize - - properties := gopter.NewProperties(parameters) - - properties.Property("l'th lagrange polynomials must evaluate to 1 on l and 0 on other values in the domain", prop.ForAll( - testForDomainSize, - gen.UInt8Range(2, maxLagrangeDomainSize), - )) - - properties.TestingRun(t, gopter.ConsoleReporter(false)) -} - -func TestLagrangeCache(t *testing.T) { - for _, i := range []int{5, 2, 8, 4, 6, 3, 0} { - b := getLagrangeBasis(uint8(i)) - assert.Equal(t, b, getLagrangeBasis(uint8(i))) // second call must yield the same result - } -} diff --git a/field/goldilocks/polynomial/pool.go b/field/goldilocks/polynomial/pool.go deleted file mode 100644 index 7ee7ee2213..0000000000 --- a/field/goldilocks/polynomial/pool.go +++ /dev/null @@ -1,191 +0,0 @@ -// Copyright 2020-2026 Consensys Software Inc. -// Licensed under the Apache License, Version 2.0. See the LICENSE file for details. - -// Code generated by consensys/gnark-crypto DO NOT EDIT - -package polynomial - -import ( - "encoding/json" - "fmt" - "runtime" - "sort" - "sync" - "unsafe" - - "github.com/consensys/gnark-crypto/field/goldilocks" -) - -// Memory management for polynomials -// WARNING: This is not thread safe TODO: Make sure that is not a problem -// TODO: There is a lot of "unsafe" memory management here and needs to be vetted thoroughly - -type sizedPool struct { - maxN int - pool sync.Pool - stats poolStats -} - -type inUseData struct { - allocatedFor []uintptr - pool *sizedPool -} - -type Pool struct { - //lock sync.Mutex - inUse sync.Map - subPools []sizedPool -} - -func (p *sizedPool) get(n int) *goldilocks.Element { - p.stats.make(n) - return p.pool.Get().(*goldilocks.Element) -} - -func (p *sizedPool) put(ptr *goldilocks.Element) { - p.stats.dump() - p.pool.Put(ptr) -} - -func NewPool(maxN ...int) (pool Pool) { - - sort.Ints(maxN) - pool = Pool{ - subPools: make([]sizedPool, len(maxN)), - } - - for i := range pool.subPools { - subPool := &pool.subPools[i] - subPool.maxN = maxN[i] - subPool.pool = sync.Pool{ - New: func() any { - subPool.stats.Allocated++ - return getDataPointer(make([]goldilocks.Element, 0, subPool.maxN)) - }, - } - } - return -} - -func (p *Pool) findCorrespondingPool(n int) *sizedPool { - poolI := 0 - for poolI < len(p.subPools) && n > p.subPools[poolI].maxN { - poolI++ - } - return &p.subPools[poolI] // out of bounds error here would mean that n is too large -} - -func (p *Pool) Make(n int) []goldilocks.Element { - pool := p.findCorrespondingPool(n) - ptr := pool.get(n) - p.addInUse(ptr, pool) - return unsafe.Slice(ptr, n) -} - -// Dump dumps a set of polynomials into the pool -func (p *Pool) Dump(slices ...[]goldilocks.Element) { - for _, slice := range slices { - ptr := getDataPointer(slice) - if metadata, ok := p.inUse.Load(ptr); ok { - p.inUse.Delete(ptr) - metadata.(inUseData).pool.put(ptr) - } else { - panic("attempting to dump a slice not created by the pool") - } - } -} - -func (p *Pool) addInUse(ptr *goldilocks.Element, pool *sizedPool) { - pcs := make([]uintptr, 2) - n := runtime.Callers(3, pcs) - - if prevPcs, ok := p.inUse.Load(ptr); ok { // TODO: remove if unnecessary for security - panic(fmt.Errorf("re-allocated non-dumped slice, previously allocated at %v", runtime.CallersFrames(prevPcs.(inUseData).allocatedFor))) - } - p.inUse.Store(ptr, inUseData{ - allocatedFor: pcs[:n], - pool: pool, - }) -} - -func printFrame(frame runtime.Frame) { - fmt.Printf("\t%s line %d, function %s\n", frame.File, frame.Line, frame.Function) -} - -func (p *Pool) printInUse() { - fmt.Println("slices never dumped allocated at:") - p.inUse.Range(func(_, pcs any) bool { - fmt.Println("-------------------------") - - var frame runtime.Frame - frames := runtime.CallersFrames(pcs.(inUseData).allocatedFor) - more := true - for more { - frame, more = frames.Next() - printFrame(frame) - } - return true - }) -} - -type poolStats struct { - Used int - Allocated int - ReuseRate float64 - InUse int - GreatestNUsed int - SmallestNUsed int -} - -type poolsStats struct { - SubPools []poolStats - InUse int -} - -func (s *poolStats) make(n int) { - s.Used++ - s.InUse++ - if n > s.GreatestNUsed { - s.GreatestNUsed = n - } - if s.SmallestNUsed == 0 || s.SmallestNUsed > n { - s.SmallestNUsed = n - } -} - -func (s *poolStats) dump() { - s.InUse-- -} - -func (s *poolStats) finalize() { - s.ReuseRate = float64(s.Used) / float64(s.Allocated) -} - -func getDataPointer(slice []goldilocks.Element) *goldilocks.Element { - return (*goldilocks.Element)(unsafe.SliceData(slice)) -} - -func (p *Pool) PrintPoolStats() { - InUse := 0 - subStats := make([]poolStats, len(p.subPools)) - for i := range p.subPools { - subPool := &p.subPools[i] - subPool.stats.finalize() - subStats[i] = subPool.stats - InUse += subPool.stats.InUse - } - - stats := poolsStats{ - SubPools: subStats, - InUse: InUse, - } - serialized, _ := json.MarshalIndent(stats, "", " ") - fmt.Println(string(serialized)) - p.printInUse() -} - -func (p *Pool) Clone(slice []goldilocks.Element) []goldilocks.Element { - res := p.Make(len(slice)) - copy(res, slice) - return res -} diff --git a/internal/generator/main.go b/internal/generator/main.go index 4e4129b879..a6751407d0 100644 --- a/internal/generator/main.go +++ b/internal/generator/main.go @@ -91,12 +91,16 @@ func main() { field.WithIOP(), )) - // polynomial package (Polynomial, MultiLin, Pool, ...) over the base field - assertNoError(polynomial.Generate(fieldConfig.FieldDependency{ - FieldPackagePath: "github.com/consensys/gnark-crypto/field/" + f.Name, - FieldPackageName: f.Name, - ElementType: f.Name + ".Element", - }, filepath.Join(outputDir, "polynomial"), true, true, gen)) + // polynomial package (Polynomial, MultiLin, Pool, ...) over the base + // field. It is only needed by the polynomial packages over the + // extensions, whose FoldFromBase takes a base field MultiLin. + if len(f.PolynomialExtensions) > 0 { + assertNoError(polynomial.Generate(fieldConfig.FieldDependency{ + FieldPackagePath: "github.com/consensys/gnark-crypto/field/" + f.Name, + FieldPackageName: f.Name, + ElementType: f.Name + ".Element", + }, filepath.Join(outputDir, "polynomial"), true, true, gen)) + } // polynomial package (Polynomial, MultiLin, Pool, ...) over the // requested extensions From 5798e085ed745f1a8e3364510768029602c2a428 Mon Sep 17 00:00:00 2001 From: Arya Tabaie Date: Wed, 7 Oct 2026 13:41:08 -0500 Subject: [PATCH 18/18] revert: move the Poseidon2 matrices to a separate PR Revert "build: go generate" and "feat: export external and internal matrices". They are unrelated to the extension work of this PR and continue in feat/poseidon2-matrices. Co-Authored-By: Claude Sonnet 5.5 Signed-off-by: Arya Tabaie --- ecc/bls12-377/fr/poseidon2/poseidon2.go | 38 ----- ecc/bls12-377/fr/poseidon2/poseidon2_test.go | 46 ------- ecc/bls12-381/fr/poseidon2/poseidon2.go | 38 ----- ecc/bls12-381/fr/poseidon2/poseidon2_test.go | 46 ------- ecc/bls24-315/fr/poseidon2/poseidon2.go | 38 ----- ecc/bls24-315/fr/poseidon2/poseidon2_test.go | 46 ------- ecc/bls24-317/fr/poseidon2/poseidon2.go | 38 ----- ecc/bls24-317/fr/poseidon2/poseidon2_test.go | 46 ------- ecc/bn254/fr/poseidon2/poseidon2.go | 41 ------ ecc/bn254/fr/poseidon2/poseidon2_test.go | 50 ------- ecc/bw6-633/fr/poseidon2/poseidon2.go | 38 ----- ecc/bw6-633/fr/poseidon2/poseidon2_test.go | 46 ------- ecc/bw6-761/fr/poseidon2/poseidon2.go | 38 ----- ecc/bw6-761/fr/poseidon2/poseidon2_test.go | 46 ------- ecc/grumpkin/fr/poseidon2/poseidon2.go | 38 ----- ecc/grumpkin/fr/poseidon2/poseidon2_test.go | 46 ------- field/babybear/poseidon2/poseidon2.go | 31 ----- field/babybear/poseidon2/poseidon2_test.go | 46 ------- field/goldilocks/poseidon2/poseidon2.go | 31 ----- field/goldilocks/poseidon2/poseidon2_test.go | 46 ------- field/koalabear/poseidon2/poseidon2.go | 31 ----- field/koalabear/poseidon2/poseidon2_test.go | 46 ------- field/mamabear/poseidon2/poseidon2.go | 31 ----- field/mamabear/poseidon2/poseidon2_test.go | 46 ------- .../hash/poseidon2/template/poseidon2.go.tmpl | 45 ------ .../poseidon2/template/poseidon2.test.go.tmpl | 54 +------- .../template/poseidon2/poseidon2.go.tmpl | 31 ----- .../template/poseidon2/poseidon2_test.go.tmpl | 47 ------- internal/poseidon2/matrices.go | 121 ---------------- internal/poseidon2/matrices_test.go | 130 ------------------ 30 files changed, 1 insertion(+), 1414 deletions(-) delete mode 100644 internal/poseidon2/matrices.go delete mode 100644 internal/poseidon2/matrices_test.go diff --git a/ecc/bls12-377/fr/poseidon2/poseidon2.go b/ecc/bls12-377/fr/poseidon2/poseidon2.go index ce9c17f44a..73558f2e92 100644 --- a/ecc/bls12-377/fr/poseidon2/poseidon2.go +++ b/ecc/bls12-377/fr/poseidon2/poseidon2.go @@ -12,7 +12,6 @@ import ( "golang.org/x/crypto/sha3" "github.com/consensys/gnark-crypto/ecc/bls12-377/fr" - internalposeidon2 "github.com/consensys/gnark-crypto/internal/poseidon2" ) var ( @@ -122,43 +121,6 @@ type Permutation struct { params *Parameters } -// ExternalMatrix returns the dense external matrix M_E of the permutation, which -// is applied to the state before the first round and after each full round. -// With I and J the identity and the all-ones matrices, and M4 the 4×4 block of -// the Poseidon2 paper (rows 5 7 1 3 / 4 6 1 1 / 1 3 5 7 / 1 1 4 6), M_E is -// -// - I + J for width 2 and 3; -// - M4 for width 4; -// - circ(2·M4, M4, ..., M4) = (I + J) ⊗ M4 for width 4k with k > 1. -func (p *Parameters) ExternalMatrix() [][]fr.Element { - return internalposeidon2.ExternalMatrix[fr.Element](p.Width, internalposeidon2.Paper) -} - -// InternalMatrix returns the dense internal matrix M_I of the permutation, which -// is applied to the state after each partial round. With J the all-ones matrix, -// -// M_I = J + diag(d), -// -// that is, for a state x, (M_I·x)_i = Σ_j x_j + d_i·x_i, where d is the diagonal -// of the permutation. -func (p *Parameters) InternalMatrix() [][]fr.Element { - var d []fr.Element - switch p.Width { - case 2: - d = make([]fr.Element, 2) - d[0].SetOne() - d[1].SetUint64(2) - case 3: - d = make([]fr.Element, 3) - d[0].SetOne() - d[1].SetOne() - d[2].SetUint64(2) - default: - panic("only Width=2,3 are supported") - } - return internalposeidon2.InternalMatrix(d) -} - // NewPermutation returns a new Poseidon2 permutation instance. func NewPermutation(t, rf, rp int) *Permutation { if t < 2 || t > 3 { diff --git a/ecc/bls12-377/fr/poseidon2/poseidon2_test.go b/ecc/bls12-377/fr/poseidon2/poseidon2_test.go index 17b7c0b70b..28647f8d12 100644 --- a/ecc/bls12-377/fr/poseidon2/poseidon2_test.go +++ b/ecc/bls12-377/fr/poseidon2/poseidon2_test.go @@ -109,49 +109,3 @@ func TestHashReset(t *testing.T) { require.Equal(t, res, h.Sum(nil)) } - -// denseMatrixMul returns m·x. -func denseMatrixMul(m [][]fr.Element, x []fr.Element) []fr.Element { - res := make([]fr.Element, len(m)) - for i := range m { - var tmp fr.Element - for j := range m[i] { - tmp.Mul(&m[i][j], &x[j]) - res[i].Add(&res[i], &tmp) - } - } - return res -} - -// TestDenseMatricesMatchInPlace checks that the dense matrices act on a random -// state as the in-place multiplications of the permutation. -func TestDenseMatricesMatchInPlace(t *testing.T) { - for _, tc := range []struct { - width, nbFullRounds, nbPartialRounds int - }{ - {2, 8, 56}, - {3, 8, 56}, - } { - h := NewPermutation(tc.width, tc.nbFullRounds, tc.nbPartialRounds) - - for name, step := range map[string]struct { - dense [][]fr.Element - inPlace func([]fr.Element) - }{ - "external": {h.params.ExternalMatrix(), h.matMulExternalInPlace}, - "internal": {h.params.InternalMatrix(), h.matMulInternalInPlace}, - } { - x := make([]fr.Element, tc.width) - for i := range x { - x[i].MustSetRandom() - } - want := denseMatrixMul(step.dense, x) - step.inPlace(x) - for i := range x { - if !x[i].Equal(&want[i]) { - t.Fatalf("width %d, %s matrix: dense and in-place multiplications differ at index %d", tc.width, name, i) - } - } - } - } -} diff --git a/ecc/bls12-381/fr/poseidon2/poseidon2.go b/ecc/bls12-381/fr/poseidon2/poseidon2.go index 70d662a45f..b98a7b4846 100644 --- a/ecc/bls12-381/fr/poseidon2/poseidon2.go +++ b/ecc/bls12-381/fr/poseidon2/poseidon2.go @@ -12,7 +12,6 @@ import ( "golang.org/x/crypto/sha3" "github.com/consensys/gnark-crypto/ecc/bls12-381/fr" - internalposeidon2 "github.com/consensys/gnark-crypto/internal/poseidon2" ) var ( @@ -122,43 +121,6 @@ type Permutation struct { params *Parameters } -// ExternalMatrix returns the dense external matrix M_E of the permutation, which -// is applied to the state before the first round and after each full round. -// With I and J the identity and the all-ones matrices, and M4 the 4×4 block of -// the Poseidon2 paper (rows 5 7 1 3 / 4 6 1 1 / 1 3 5 7 / 1 1 4 6), M_E is -// -// - I + J for width 2 and 3; -// - M4 for width 4; -// - circ(2·M4, M4, ..., M4) = (I + J) ⊗ M4 for width 4k with k > 1. -func (p *Parameters) ExternalMatrix() [][]fr.Element { - return internalposeidon2.ExternalMatrix[fr.Element](p.Width, internalposeidon2.Paper) -} - -// InternalMatrix returns the dense internal matrix M_I of the permutation, which -// is applied to the state after each partial round. With J the all-ones matrix, -// -// M_I = J + diag(d), -// -// that is, for a state x, (M_I·x)_i = Σ_j x_j + d_i·x_i, where d is the diagonal -// of the permutation. -func (p *Parameters) InternalMatrix() [][]fr.Element { - var d []fr.Element - switch p.Width { - case 2: - d = make([]fr.Element, 2) - d[0].SetOne() - d[1].SetUint64(2) - case 3: - d = make([]fr.Element, 3) - d[0].SetOne() - d[1].SetOne() - d[2].SetUint64(2) - default: - panic("only Width=2,3 are supported") - } - return internalposeidon2.InternalMatrix(d) -} - // NewPermutation returns a new Poseidon2 permutation instance. func NewPermutation(t, rf, rp int) *Permutation { if t < 2 || t > 3 { diff --git a/ecc/bls12-381/fr/poseidon2/poseidon2_test.go b/ecc/bls12-381/fr/poseidon2/poseidon2_test.go index 324edceef5..c577df2f7c 100644 --- a/ecc/bls12-381/fr/poseidon2/poseidon2_test.go +++ b/ecc/bls12-381/fr/poseidon2/poseidon2_test.go @@ -109,49 +109,3 @@ func TestHashReset(t *testing.T) { require.Equal(t, res, h.Sum(nil)) } - -// denseMatrixMul returns m·x. -func denseMatrixMul(m [][]fr.Element, x []fr.Element) []fr.Element { - res := make([]fr.Element, len(m)) - for i := range m { - var tmp fr.Element - for j := range m[i] { - tmp.Mul(&m[i][j], &x[j]) - res[i].Add(&res[i], &tmp) - } - } - return res -} - -// TestDenseMatricesMatchInPlace checks that the dense matrices act on a random -// state as the in-place multiplications of the permutation. -func TestDenseMatricesMatchInPlace(t *testing.T) { - for _, tc := range []struct { - width, nbFullRounds, nbPartialRounds int - }{ - {2, 8, 56}, - {3, 8, 56}, - } { - h := NewPermutation(tc.width, tc.nbFullRounds, tc.nbPartialRounds) - - for name, step := range map[string]struct { - dense [][]fr.Element - inPlace func([]fr.Element) - }{ - "external": {h.params.ExternalMatrix(), h.matMulExternalInPlace}, - "internal": {h.params.InternalMatrix(), h.matMulInternalInPlace}, - } { - x := make([]fr.Element, tc.width) - for i := range x { - x[i].MustSetRandom() - } - want := denseMatrixMul(step.dense, x) - step.inPlace(x) - for i := range x { - if !x[i].Equal(&want[i]) { - t.Fatalf("width %d, %s matrix: dense and in-place multiplications differ at index %d", tc.width, name, i) - } - } - } - } -} diff --git a/ecc/bls24-315/fr/poseidon2/poseidon2.go b/ecc/bls24-315/fr/poseidon2/poseidon2.go index 6d233573c3..6c4be772a2 100644 --- a/ecc/bls24-315/fr/poseidon2/poseidon2.go +++ b/ecc/bls24-315/fr/poseidon2/poseidon2.go @@ -12,7 +12,6 @@ import ( "golang.org/x/crypto/sha3" "github.com/consensys/gnark-crypto/ecc/bls24-315/fr" - internalposeidon2 "github.com/consensys/gnark-crypto/internal/poseidon2" ) var ( @@ -122,43 +121,6 @@ type Permutation struct { params *Parameters } -// ExternalMatrix returns the dense external matrix M_E of the permutation, which -// is applied to the state before the first round and after each full round. -// With I and J the identity and the all-ones matrices, and M4 the 4×4 block of -// the Poseidon2 paper (rows 5 7 1 3 / 4 6 1 1 / 1 3 5 7 / 1 1 4 6), M_E is -// -// - I + J for width 2 and 3; -// - M4 for width 4; -// - circ(2·M4, M4, ..., M4) = (I + J) ⊗ M4 for width 4k with k > 1. -func (p *Parameters) ExternalMatrix() [][]fr.Element { - return internalposeidon2.ExternalMatrix[fr.Element](p.Width, internalposeidon2.Paper) -} - -// InternalMatrix returns the dense internal matrix M_I of the permutation, which -// is applied to the state after each partial round. With J the all-ones matrix, -// -// M_I = J + diag(d), -// -// that is, for a state x, (M_I·x)_i = Σ_j x_j + d_i·x_i, where d is the diagonal -// of the permutation. -func (p *Parameters) InternalMatrix() [][]fr.Element { - var d []fr.Element - switch p.Width { - case 2: - d = make([]fr.Element, 2) - d[0].SetOne() - d[1].SetUint64(2) - case 3: - d = make([]fr.Element, 3) - d[0].SetOne() - d[1].SetOne() - d[2].SetUint64(2) - default: - panic("only Width=2,3 are supported") - } - return internalposeidon2.InternalMatrix(d) -} - // NewPermutation returns a new Poseidon2 permutation instance. func NewPermutation(t, rf, rp int) *Permutation { if t < 2 || t > 3 { diff --git a/ecc/bls24-315/fr/poseidon2/poseidon2_test.go b/ecc/bls24-315/fr/poseidon2/poseidon2_test.go index e35d355476..1a2f1daedc 100644 --- a/ecc/bls24-315/fr/poseidon2/poseidon2_test.go +++ b/ecc/bls24-315/fr/poseidon2/poseidon2_test.go @@ -109,49 +109,3 @@ func TestHashReset(t *testing.T) { require.Equal(t, res, h.Sum(nil)) } - -// denseMatrixMul returns m·x. -func denseMatrixMul(m [][]fr.Element, x []fr.Element) []fr.Element { - res := make([]fr.Element, len(m)) - for i := range m { - var tmp fr.Element - for j := range m[i] { - tmp.Mul(&m[i][j], &x[j]) - res[i].Add(&res[i], &tmp) - } - } - return res -} - -// TestDenseMatricesMatchInPlace checks that the dense matrices act on a random -// state as the in-place multiplications of the permutation. -func TestDenseMatricesMatchInPlace(t *testing.T) { - for _, tc := range []struct { - width, nbFullRounds, nbPartialRounds int - }{ - {2, 8, 56}, - {3, 8, 56}, - } { - h := NewPermutation(tc.width, tc.nbFullRounds, tc.nbPartialRounds) - - for name, step := range map[string]struct { - dense [][]fr.Element - inPlace func([]fr.Element) - }{ - "external": {h.params.ExternalMatrix(), h.matMulExternalInPlace}, - "internal": {h.params.InternalMatrix(), h.matMulInternalInPlace}, - } { - x := make([]fr.Element, tc.width) - for i := range x { - x[i].MustSetRandom() - } - want := denseMatrixMul(step.dense, x) - step.inPlace(x) - for i := range x { - if !x[i].Equal(&want[i]) { - t.Fatalf("width %d, %s matrix: dense and in-place multiplications differ at index %d", tc.width, name, i) - } - } - } - } -} diff --git a/ecc/bls24-317/fr/poseidon2/poseidon2.go b/ecc/bls24-317/fr/poseidon2/poseidon2.go index e1809d2f86..7e1061a638 100644 --- a/ecc/bls24-317/fr/poseidon2/poseidon2.go +++ b/ecc/bls24-317/fr/poseidon2/poseidon2.go @@ -12,7 +12,6 @@ import ( "golang.org/x/crypto/sha3" "github.com/consensys/gnark-crypto/ecc/bls24-317/fr" - internalposeidon2 "github.com/consensys/gnark-crypto/internal/poseidon2" ) var ( @@ -122,43 +121,6 @@ type Permutation struct { params *Parameters } -// ExternalMatrix returns the dense external matrix M_E of the permutation, which -// is applied to the state before the first round and after each full round. -// With I and J the identity and the all-ones matrices, and M4 the 4×4 block of -// the Poseidon2 paper (rows 5 7 1 3 / 4 6 1 1 / 1 3 5 7 / 1 1 4 6), M_E is -// -// - I + J for width 2 and 3; -// - M4 for width 4; -// - circ(2·M4, M4, ..., M4) = (I + J) ⊗ M4 for width 4k with k > 1. -func (p *Parameters) ExternalMatrix() [][]fr.Element { - return internalposeidon2.ExternalMatrix[fr.Element](p.Width, internalposeidon2.Paper) -} - -// InternalMatrix returns the dense internal matrix M_I of the permutation, which -// is applied to the state after each partial round. With J the all-ones matrix, -// -// M_I = J + diag(d), -// -// that is, for a state x, (M_I·x)_i = Σ_j x_j + d_i·x_i, where d is the diagonal -// of the permutation. -func (p *Parameters) InternalMatrix() [][]fr.Element { - var d []fr.Element - switch p.Width { - case 2: - d = make([]fr.Element, 2) - d[0].SetOne() - d[1].SetUint64(2) - case 3: - d = make([]fr.Element, 3) - d[0].SetOne() - d[1].SetOne() - d[2].SetUint64(2) - default: - panic("only Width=2,3 are supported") - } - return internalposeidon2.InternalMatrix(d) -} - // NewPermutation returns a new Poseidon2 permutation instance. func NewPermutation(t, rf, rp int) *Permutation { if t < 2 || t > 3 { diff --git a/ecc/bls24-317/fr/poseidon2/poseidon2_test.go b/ecc/bls24-317/fr/poseidon2/poseidon2_test.go index 1165a8ca8b..b89e8741f5 100644 --- a/ecc/bls24-317/fr/poseidon2/poseidon2_test.go +++ b/ecc/bls24-317/fr/poseidon2/poseidon2_test.go @@ -109,49 +109,3 @@ func TestHashReset(t *testing.T) { require.Equal(t, res, h.Sum(nil)) } - -// denseMatrixMul returns m·x. -func denseMatrixMul(m [][]fr.Element, x []fr.Element) []fr.Element { - res := make([]fr.Element, len(m)) - for i := range m { - var tmp fr.Element - for j := range m[i] { - tmp.Mul(&m[i][j], &x[j]) - res[i].Add(&res[i], &tmp) - } - } - return res -} - -// TestDenseMatricesMatchInPlace checks that the dense matrices act on a random -// state as the in-place multiplications of the permutation. -func TestDenseMatricesMatchInPlace(t *testing.T) { - for _, tc := range []struct { - width, nbFullRounds, nbPartialRounds int - }{ - {2, 8, 56}, - {3, 8, 56}, - } { - h := NewPermutation(tc.width, tc.nbFullRounds, tc.nbPartialRounds) - - for name, step := range map[string]struct { - dense [][]fr.Element - inPlace func([]fr.Element) - }{ - "external": {h.params.ExternalMatrix(), h.matMulExternalInPlace}, - "internal": {h.params.InternalMatrix(), h.matMulInternalInPlace}, - } { - x := make([]fr.Element, tc.width) - for i := range x { - x[i].MustSetRandom() - } - want := denseMatrixMul(step.dense, x) - step.inPlace(x) - for i := range x { - if !x[i].Equal(&want[i]) { - t.Fatalf("width %d, %s matrix: dense and in-place multiplications differ at index %d", tc.width, name, i) - } - } - } - } -} diff --git a/ecc/bn254/fr/poseidon2/poseidon2.go b/ecc/bn254/fr/poseidon2/poseidon2.go index 09ab10e984..66ca9e8830 100644 --- a/ecc/bn254/fr/poseidon2/poseidon2.go +++ b/ecc/bn254/fr/poseidon2/poseidon2.go @@ -12,7 +12,6 @@ import ( "golang.org/x/crypto/sha3" "github.com/consensys/gnark-crypto/ecc/bn254/fr" - internalposeidon2 "github.com/consensys/gnark-crypto/internal/poseidon2" ) var ( @@ -141,46 +140,6 @@ type Permutation struct { params *Parameters } -// ExternalMatrix returns the dense external matrix M_E of the permutation, which -// is applied to the state before the first round and after each full round. -// With I and J the identity and the all-ones matrices, and M4 the 4×4 block of -// the Poseidon2 paper (rows 5 7 1 3 / 4 6 1 1 / 1 3 5 7 / 1 1 4 6), M_E is -// -// - I + J for width 2 and 3; -// - M4 for width 4; -// - circ(2·M4, M4, ..., M4) = (I + J) ⊗ M4 for width 4k with k > 1. -func (p *Parameters) ExternalMatrix() [][]fr.Element { - return internalposeidon2.ExternalMatrix[fr.Element](p.Width, internalposeidon2.Paper) -} - -// InternalMatrix returns the dense internal matrix M_I of the permutation, which -// is applied to the state after each partial round. With J the all-ones matrix, -// -// M_I = J + diag(d), -// -// that is, for a state x, (M_I·x)_i = Σ_j x_j + d_i·x_i, where d is the diagonal -// of the permutation (DiagM1 for the widths of 4 and more). -func (p *Parameters) InternalMatrix() [][]fr.Element { - var d []fr.Element - switch p.Width { - case 2: - d = make([]fr.Element, 2) - d[0].SetOne() - d[1].SetUint64(2) - case 3: - d = make([]fr.Element, 3) - d[0].SetOne() - d[1].SetOne() - d[2].SetUint64(2) - default: - if len(p.DiagM1) != p.Width { - panic("missing internal matrix diagonal") - } - d = p.DiagM1 - } - return internalposeidon2.InternalMatrix(d) -} - // NewPermutation returns a new Poseidon2 permutation instance. func NewPermutation(t, rf, rp int) *Permutation { if t == 2 || t == 3 { diff --git a/ecc/bn254/fr/poseidon2/poseidon2_test.go b/ecc/bn254/fr/poseidon2/poseidon2_test.go index e672ca94fa..9990070916 100644 --- a/ecc/bn254/fr/poseidon2/poseidon2_test.go +++ b/ecc/bn254/fr/poseidon2/poseidon2_test.go @@ -193,53 +193,3 @@ func TestHashReset(t *testing.T) { require.Equal(t, res, h.Sum(nil)) } - -// denseMatrixMul returns m·x. -func denseMatrixMul(m [][]fr.Element, x []fr.Element) []fr.Element { - res := make([]fr.Element, len(m)) - for i := range m { - var tmp fr.Element - for j := range m[i] { - tmp.Mul(&m[i][j], &x[j]) - res[i].Add(&res[i], &tmp) - } - } - return res -} - -// TestDenseMatricesMatchInPlace checks that the dense matrices act on a random -// state as the in-place multiplications of the permutation. -func TestDenseMatricesMatchInPlace(t *testing.T) { - for _, tc := range []struct { - width, nbFullRounds, nbPartialRounds int - }{ - {2, 8, 56}, - {3, 8, 56}, - {4, 8, 56}, - {8, 8, 57}, - {12, 8, 57}, - {16, 8, 57}, - } { - h := NewPermutation(tc.width, tc.nbFullRounds, tc.nbPartialRounds) - - for name, step := range map[string]struct { - dense [][]fr.Element - inPlace func([]fr.Element) - }{ - "external": {h.params.ExternalMatrix(), h.matMulExternalInPlace}, - "internal": {h.params.InternalMatrix(), h.matMulInternalInPlace}, - } { - x := make([]fr.Element, tc.width) - for i := range x { - x[i].MustSetRandom() - } - want := denseMatrixMul(step.dense, x) - step.inPlace(x) - for i := range x { - if !x[i].Equal(&want[i]) { - t.Fatalf("width %d, %s matrix: dense and in-place multiplications differ at index %d", tc.width, name, i) - } - } - } - } -} diff --git a/ecc/bw6-633/fr/poseidon2/poseidon2.go b/ecc/bw6-633/fr/poseidon2/poseidon2.go index 35b47b266f..e577b9dd88 100644 --- a/ecc/bw6-633/fr/poseidon2/poseidon2.go +++ b/ecc/bw6-633/fr/poseidon2/poseidon2.go @@ -12,7 +12,6 @@ import ( "golang.org/x/crypto/sha3" "github.com/consensys/gnark-crypto/ecc/bw6-633/fr" - internalposeidon2 "github.com/consensys/gnark-crypto/internal/poseidon2" ) var ( @@ -122,43 +121,6 @@ type Permutation struct { params *Parameters } -// ExternalMatrix returns the dense external matrix M_E of the permutation, which -// is applied to the state before the first round and after each full round. -// With I and J the identity and the all-ones matrices, and M4 the 4×4 block of -// the Poseidon2 paper (rows 5 7 1 3 / 4 6 1 1 / 1 3 5 7 / 1 1 4 6), M_E is -// -// - I + J for width 2 and 3; -// - M4 for width 4; -// - circ(2·M4, M4, ..., M4) = (I + J) ⊗ M4 for width 4k with k > 1. -func (p *Parameters) ExternalMatrix() [][]fr.Element { - return internalposeidon2.ExternalMatrix[fr.Element](p.Width, internalposeidon2.Paper) -} - -// InternalMatrix returns the dense internal matrix M_I of the permutation, which -// is applied to the state after each partial round. With J the all-ones matrix, -// -// M_I = J + diag(d), -// -// that is, for a state x, (M_I·x)_i = Σ_j x_j + d_i·x_i, where d is the diagonal -// of the permutation. -func (p *Parameters) InternalMatrix() [][]fr.Element { - var d []fr.Element - switch p.Width { - case 2: - d = make([]fr.Element, 2) - d[0].SetOne() - d[1].SetUint64(2) - case 3: - d = make([]fr.Element, 3) - d[0].SetOne() - d[1].SetOne() - d[2].SetUint64(2) - default: - panic("only Width=2,3 are supported") - } - return internalposeidon2.InternalMatrix(d) -} - // NewPermutation returns a new Poseidon2 permutation instance. func NewPermutation(t, rf, rp int) *Permutation { if t < 2 || t > 3 { diff --git a/ecc/bw6-633/fr/poseidon2/poseidon2_test.go b/ecc/bw6-633/fr/poseidon2/poseidon2_test.go index 4bb1318954..06de110947 100644 --- a/ecc/bw6-633/fr/poseidon2/poseidon2_test.go +++ b/ecc/bw6-633/fr/poseidon2/poseidon2_test.go @@ -109,49 +109,3 @@ func TestHashReset(t *testing.T) { require.Equal(t, res, h.Sum(nil)) } - -// denseMatrixMul returns m·x. -func denseMatrixMul(m [][]fr.Element, x []fr.Element) []fr.Element { - res := make([]fr.Element, len(m)) - for i := range m { - var tmp fr.Element - for j := range m[i] { - tmp.Mul(&m[i][j], &x[j]) - res[i].Add(&res[i], &tmp) - } - } - return res -} - -// TestDenseMatricesMatchInPlace checks that the dense matrices act on a random -// state as the in-place multiplications of the permutation. -func TestDenseMatricesMatchInPlace(t *testing.T) { - for _, tc := range []struct { - width, nbFullRounds, nbPartialRounds int - }{ - {2, 8, 56}, - {3, 8, 56}, - } { - h := NewPermutation(tc.width, tc.nbFullRounds, tc.nbPartialRounds) - - for name, step := range map[string]struct { - dense [][]fr.Element - inPlace func([]fr.Element) - }{ - "external": {h.params.ExternalMatrix(), h.matMulExternalInPlace}, - "internal": {h.params.InternalMatrix(), h.matMulInternalInPlace}, - } { - x := make([]fr.Element, tc.width) - for i := range x { - x[i].MustSetRandom() - } - want := denseMatrixMul(step.dense, x) - step.inPlace(x) - for i := range x { - if !x[i].Equal(&want[i]) { - t.Fatalf("width %d, %s matrix: dense and in-place multiplications differ at index %d", tc.width, name, i) - } - } - } - } -} diff --git a/ecc/bw6-761/fr/poseidon2/poseidon2.go b/ecc/bw6-761/fr/poseidon2/poseidon2.go index 8ae5a83a66..ee908ff421 100644 --- a/ecc/bw6-761/fr/poseidon2/poseidon2.go +++ b/ecc/bw6-761/fr/poseidon2/poseidon2.go @@ -12,7 +12,6 @@ import ( "golang.org/x/crypto/sha3" "github.com/consensys/gnark-crypto/ecc/bw6-761/fr" - internalposeidon2 "github.com/consensys/gnark-crypto/internal/poseidon2" ) var ( @@ -122,43 +121,6 @@ type Permutation struct { params *Parameters } -// ExternalMatrix returns the dense external matrix M_E of the permutation, which -// is applied to the state before the first round and after each full round. -// With I and J the identity and the all-ones matrices, and M4 the 4×4 block of -// the Poseidon2 paper (rows 5 7 1 3 / 4 6 1 1 / 1 3 5 7 / 1 1 4 6), M_E is -// -// - I + J for width 2 and 3; -// - M4 for width 4; -// - circ(2·M4, M4, ..., M4) = (I + J) ⊗ M4 for width 4k with k > 1. -func (p *Parameters) ExternalMatrix() [][]fr.Element { - return internalposeidon2.ExternalMatrix[fr.Element](p.Width, internalposeidon2.Paper) -} - -// InternalMatrix returns the dense internal matrix M_I of the permutation, which -// is applied to the state after each partial round. With J the all-ones matrix, -// -// M_I = J + diag(d), -// -// that is, for a state x, (M_I·x)_i = Σ_j x_j + d_i·x_i, where d is the diagonal -// of the permutation. -func (p *Parameters) InternalMatrix() [][]fr.Element { - var d []fr.Element - switch p.Width { - case 2: - d = make([]fr.Element, 2) - d[0].SetOne() - d[1].SetUint64(2) - case 3: - d = make([]fr.Element, 3) - d[0].SetOne() - d[1].SetOne() - d[2].SetUint64(2) - default: - panic("only Width=2,3 are supported") - } - return internalposeidon2.InternalMatrix(d) -} - // NewPermutation returns a new Poseidon2 permutation instance. func NewPermutation(t, rf, rp int) *Permutation { if t < 2 || t > 3 { diff --git a/ecc/bw6-761/fr/poseidon2/poseidon2_test.go b/ecc/bw6-761/fr/poseidon2/poseidon2_test.go index 5264feb873..9e3a52cd7c 100644 --- a/ecc/bw6-761/fr/poseidon2/poseidon2_test.go +++ b/ecc/bw6-761/fr/poseidon2/poseidon2_test.go @@ -109,49 +109,3 @@ func TestHashReset(t *testing.T) { require.Equal(t, res, h.Sum(nil)) } - -// denseMatrixMul returns m·x. -func denseMatrixMul(m [][]fr.Element, x []fr.Element) []fr.Element { - res := make([]fr.Element, len(m)) - for i := range m { - var tmp fr.Element - for j := range m[i] { - tmp.Mul(&m[i][j], &x[j]) - res[i].Add(&res[i], &tmp) - } - } - return res -} - -// TestDenseMatricesMatchInPlace checks that the dense matrices act on a random -// state as the in-place multiplications of the permutation. -func TestDenseMatricesMatchInPlace(t *testing.T) { - for _, tc := range []struct { - width, nbFullRounds, nbPartialRounds int - }{ - {2, 8, 56}, - {3, 8, 56}, - } { - h := NewPermutation(tc.width, tc.nbFullRounds, tc.nbPartialRounds) - - for name, step := range map[string]struct { - dense [][]fr.Element - inPlace func([]fr.Element) - }{ - "external": {h.params.ExternalMatrix(), h.matMulExternalInPlace}, - "internal": {h.params.InternalMatrix(), h.matMulInternalInPlace}, - } { - x := make([]fr.Element, tc.width) - for i := range x { - x[i].MustSetRandom() - } - want := denseMatrixMul(step.dense, x) - step.inPlace(x) - for i := range x { - if !x[i].Equal(&want[i]) { - t.Fatalf("width %d, %s matrix: dense and in-place multiplications differ at index %d", tc.width, name, i) - } - } - } - } -} diff --git a/ecc/grumpkin/fr/poseidon2/poseidon2.go b/ecc/grumpkin/fr/poseidon2/poseidon2.go index aba28edc08..2aaeef27fa 100644 --- a/ecc/grumpkin/fr/poseidon2/poseidon2.go +++ b/ecc/grumpkin/fr/poseidon2/poseidon2.go @@ -12,7 +12,6 @@ import ( "golang.org/x/crypto/sha3" "github.com/consensys/gnark-crypto/ecc/grumpkin/fr" - internalposeidon2 "github.com/consensys/gnark-crypto/internal/poseidon2" ) var ( @@ -122,43 +121,6 @@ type Permutation struct { params *Parameters } -// ExternalMatrix returns the dense external matrix M_E of the permutation, which -// is applied to the state before the first round and after each full round. -// With I and J the identity and the all-ones matrices, and M4 the 4×4 block of -// the Poseidon2 paper (rows 5 7 1 3 / 4 6 1 1 / 1 3 5 7 / 1 1 4 6), M_E is -// -// - I + J for width 2 and 3; -// - M4 for width 4; -// - circ(2·M4, M4, ..., M4) = (I + J) ⊗ M4 for width 4k with k > 1. -func (p *Parameters) ExternalMatrix() [][]fr.Element { - return internalposeidon2.ExternalMatrix[fr.Element](p.Width, internalposeidon2.Paper) -} - -// InternalMatrix returns the dense internal matrix M_I of the permutation, which -// is applied to the state after each partial round. With J the all-ones matrix, -// -// M_I = J + diag(d), -// -// that is, for a state x, (M_I·x)_i = Σ_j x_j + d_i·x_i, where d is the diagonal -// of the permutation. -func (p *Parameters) InternalMatrix() [][]fr.Element { - var d []fr.Element - switch p.Width { - case 2: - d = make([]fr.Element, 2) - d[0].SetOne() - d[1].SetUint64(2) - case 3: - d = make([]fr.Element, 3) - d[0].SetOne() - d[1].SetOne() - d[2].SetUint64(2) - default: - panic("only Width=2,3 are supported") - } - return internalposeidon2.InternalMatrix(d) -} - // NewPermutation returns a new Poseidon2 permutation instance. func NewPermutation(t, rf, rp int) *Permutation { if t < 2 || t > 3 { diff --git a/ecc/grumpkin/fr/poseidon2/poseidon2_test.go b/ecc/grumpkin/fr/poseidon2/poseidon2_test.go index 653198594d..d9e6b89c28 100644 --- a/ecc/grumpkin/fr/poseidon2/poseidon2_test.go +++ b/ecc/grumpkin/fr/poseidon2/poseidon2_test.go @@ -109,49 +109,3 @@ func TestHashReset(t *testing.T) { require.Equal(t, res, h.Sum(nil)) } - -// denseMatrixMul returns m·x. -func denseMatrixMul(m [][]fr.Element, x []fr.Element) []fr.Element { - res := make([]fr.Element, len(m)) - for i := range m { - var tmp fr.Element - for j := range m[i] { - tmp.Mul(&m[i][j], &x[j]) - res[i].Add(&res[i], &tmp) - } - } - return res -} - -// TestDenseMatricesMatchInPlace checks that the dense matrices act on a random -// state as the in-place multiplications of the permutation. -func TestDenseMatricesMatchInPlace(t *testing.T) { - for _, tc := range []struct { - width, nbFullRounds, nbPartialRounds int - }{ - {2, 8, 56}, - {3, 8, 56}, - } { - h := NewPermutation(tc.width, tc.nbFullRounds, tc.nbPartialRounds) - - for name, step := range map[string]struct { - dense [][]fr.Element - inPlace func([]fr.Element) - }{ - "external": {h.params.ExternalMatrix(), h.matMulExternalInPlace}, - "internal": {h.params.InternalMatrix(), h.matMulInternalInPlace}, - } { - x := make([]fr.Element, tc.width) - for i := range x { - x[i].MustSetRandom() - } - want := denseMatrixMul(step.dense, x) - step.inPlace(x) - for i := range x { - if !x[i].Equal(&want[i]) { - t.Fatalf("width %d, %s matrix: dense and in-place multiplications differ at index %d", tc.width, name, i) - } - } - } - } -} diff --git a/field/babybear/poseidon2/poseidon2.go b/field/babybear/poseidon2/poseidon2.go index 2a3b1a18f8..d303d6e414 100644 --- a/field/babybear/poseidon2/poseidon2.go +++ b/field/babybear/poseidon2/poseidon2.go @@ -16,7 +16,6 @@ import ( "golang.org/x/crypto/sha3" fr "github.com/consensys/gnark-crypto/field/babybear" - internalposeidon2 "github.com/consensys/gnark-crypto/internal/poseidon2" "github.com/consensys/gnark-crypto/utils/cpu" ) @@ -131,36 +130,6 @@ type Permutation struct { params *Parameters } -// ExternalMatrix returns the dense external matrix M_E of the permutation, which -// is applied to the state before the first round and after each full round. -// With I and J the identity and the all-ones matrices, and M4 the 4×4 block of -// Plonky3 (rows 2 3 1 1 / 1 2 3 1 / 1 1 2 3 / 3 1 1 2), M_E is -// -// - I + J for width 2 and 3; -// - M4 for width 4; -// - circ(2·M4, M4, ..., M4) = (I + J) ⊗ M4 for width 4k with k > 1. -func (p *Parameters) ExternalMatrix() [][]fr.Element { - return internalposeidon2.ExternalMatrix[fr.Element](p.Width, internalposeidon2.Plonky3) -} - -// InternalMatrix returns the dense internal matrix M_I of the permutation, which -// is applied to the state after each partial round. With J the all-ones matrix, -// -// M_I = J + diag(d), -// -// that is, for a state x, (M_I·x)_i = Σ_j x_j + d_i·x_i, where d is the diagonal -// of the permutation (diag16 for width 16, diag24 for width 24). -func (p *Parameters) InternalMatrix() [][]fr.Element { - switch p.Width { - case 16: - return internalposeidon2.InternalMatrix(diag16) - case 24: - return internalposeidon2.InternalMatrix(diag24) - default: - panic("only Width=16,24 are supported") - } -} - // NewPermutation returns a new Poseidon2 permutation instance. func NewPermutation(t, rf, rp int) *Permutation { if t != 16 && t != 24 { diff --git a/field/babybear/poseidon2/poseidon2_test.go b/field/babybear/poseidon2/poseidon2_test.go index 5317b63e7e..e2037cb4a3 100644 --- a/field/babybear/poseidon2/poseidon2_test.go +++ b/field/babybear/poseidon2/poseidon2_test.go @@ -63,52 +63,6 @@ func TestMulMulInternalInPlaceWidth24(t *testing.T) { } } -// denseMatrixMul returns m·x. -func denseMatrixMul(m [][]fr.Element, x []fr.Element) []fr.Element { - res := make([]fr.Element, len(m)) - for i := range m { - var tmp fr.Element - for j := range m[i] { - tmp.Mul(&m[i][j], &x[j]) - res[i].Add(&res[i], &tmp) - } - } - return res -} - -// TestDenseMatricesMatchInPlace checks that the dense matrices act on a random -// state as the in-place multiplications of the permutation. -func TestDenseMatricesMatchInPlace(t *testing.T) { - for _, tc := range []struct { - width, nbFullRounds, nbPartialRounds int - }{ - {16, 8, 13}, - {24, 8, 21}, - } { - h := NewPermutation(tc.width, tc.nbFullRounds, tc.nbPartialRounds) - - for name, step := range map[string]struct { - dense [][]fr.Element - inPlace func([]fr.Element) - }{ - "external": {h.params.ExternalMatrix(), h.matMulExternalInPlace}, - "internal": {h.params.InternalMatrix(), h.matMulInternalInPlace}, - } { - x := make([]fr.Element, tc.width) - for i := range x { - x[i].MustSetRandom() - } - want := denseMatrixMul(step.dense, x) - step.inPlace(x) - for i := range x { - if !x[i].Equal(&want[i]) { - t.Fatalf("width %d, %s matrix: dense and in-place multiplications differ at index %d", tc.width, name, i) - } - } - } - } -} - func TestAVX512Width16(t *testing.T) { if !cpu.SupportAVX512 { t.Skip("AVX512 not supported") diff --git a/field/goldilocks/poseidon2/poseidon2.go b/field/goldilocks/poseidon2/poseidon2.go index b76334077d..dbd15b2a56 100644 --- a/field/goldilocks/poseidon2/poseidon2.go +++ b/field/goldilocks/poseidon2/poseidon2.go @@ -15,7 +15,6 @@ import ( "golang.org/x/crypto/sha3" fr "github.com/consensys/gnark-crypto/field/goldilocks" - internalposeidon2 "github.com/consensys/gnark-crypto/internal/poseidon2" ) var ( @@ -122,36 +121,6 @@ type Permutation struct { params *Parameters } -// ExternalMatrix returns the dense external matrix M_E of the permutation, which -// is applied to the state before the first round and after each full round. -// With I and J the identity and the all-ones matrices, and M4 the 4×4 block of -// the Poseidon2 paper (rows 5 7 1 3 / 4 6 1 1 / 1 3 5 7 / 1 1 4 6), M_E is -// -// - I + J for width 2 and 3; -// - M4 for width 4; -// - circ(2·M4, M4, ..., M4) = (I + J) ⊗ M4 for width 4k with k > 1. -func (p *Parameters) ExternalMatrix() [][]fr.Element { - return internalposeidon2.ExternalMatrix[fr.Element](p.Width, internalposeidon2.Paper) -} - -// InternalMatrix returns the dense internal matrix M_I of the permutation, which -// is applied to the state after each partial round. With J the all-ones matrix, -// -// M_I = J + diag(d), -// -// that is, for a state x, (M_I·x)_i = Σ_j x_j + d_i·x_i, where d is the diagonal -// of the permutation (diag8 for width 8, diag12 for width 12). -func (p *Parameters) InternalMatrix() [][]fr.Element { - switch p.Width { - case 8: - return internalposeidon2.InternalMatrix(diag8) - case 12: - return internalposeidon2.InternalMatrix(diag12) - default: - panic("only Width=8,12 are supported") - } -} - // NewPermutation returns a new Poseidon2 permutation instance. func NewPermutation(t, rf, rp int) *Permutation { if t != 8 && t != 12 { diff --git a/field/goldilocks/poseidon2/poseidon2_test.go b/field/goldilocks/poseidon2/poseidon2_test.go index a4be7c58a0..1dfe183054 100644 --- a/field/goldilocks/poseidon2/poseidon2_test.go +++ b/field/goldilocks/poseidon2/poseidon2_test.go @@ -60,52 +60,6 @@ func TestMulMulInternalInPlaceWidth12(t *testing.T) { } } } - -// denseMatrixMul returns m·x. -func denseMatrixMul(m [][]fr.Element, x []fr.Element) []fr.Element { - res := make([]fr.Element, len(m)) - for i := range m { - var tmp fr.Element - for j := range m[i] { - tmp.Mul(&m[i][j], &x[j]) - res[i].Add(&res[i], &tmp) - } - } - return res -} - -// TestDenseMatricesMatchInPlace checks that the dense matrices act on a random -// state as the in-place multiplications of the permutation. -func TestDenseMatricesMatchInPlace(t *testing.T) { - for _, tc := range []struct { - width, nbFullRounds, nbPartialRounds int - }{ - {8, 6, 17}, - {12, 6, 17}, - } { - h := NewPermutation(tc.width, tc.nbFullRounds, tc.nbPartialRounds) - - for name, step := range map[string]struct { - dense [][]fr.Element - inPlace func([]fr.Element) - }{ - "external": {h.params.ExternalMatrix(), h.matMulExternalInPlace}, - "internal": {h.params.InternalMatrix(), h.matMulInternalInPlace}, - } { - x := make([]fr.Element, tc.width) - for i := range x { - x[i].MustSetRandom() - } - want := denseMatrixMul(step.dense, x) - step.inPlace(x) - for i := range x { - if !x[i].Equal(&want[i]) { - t.Fatalf("width %d, %s matrix: dense and in-place multiplications differ at index %d", tc.width, name, i) - } - } - } - } -} func TestPoseidon2Width8(t *testing.T) { var input, expected [8]fr.Element // these are random values generated by MustSetRandom() diff --git a/field/koalabear/poseidon2/poseidon2.go b/field/koalabear/poseidon2/poseidon2.go index c6f79cd547..abaa52da5c 100644 --- a/field/koalabear/poseidon2/poseidon2.go +++ b/field/koalabear/poseidon2/poseidon2.go @@ -16,7 +16,6 @@ import ( "golang.org/x/crypto/sha3" fr "github.com/consensys/gnark-crypto/field/koalabear" - internalposeidon2 "github.com/consensys/gnark-crypto/internal/poseidon2" "github.com/consensys/gnark-crypto/utils/cpu" ) @@ -131,36 +130,6 @@ type Permutation struct { params *Parameters } -// ExternalMatrix returns the dense external matrix M_E of the permutation, which -// is applied to the state before the first round and after each full round. -// With I and J the identity and the all-ones matrices, and M4 the 4×4 block of -// Plonky3 (rows 2 3 1 1 / 1 2 3 1 / 1 1 2 3 / 3 1 1 2), M_E is -// -// - I + J for width 2 and 3; -// - M4 for width 4; -// - circ(2·M4, M4, ..., M4) = (I + J) ⊗ M4 for width 4k with k > 1. -func (p *Parameters) ExternalMatrix() [][]fr.Element { - return internalposeidon2.ExternalMatrix[fr.Element](p.Width, internalposeidon2.Plonky3) -} - -// InternalMatrix returns the dense internal matrix M_I of the permutation, which -// is applied to the state after each partial round. With J the all-ones matrix, -// -// M_I = J + diag(d), -// -// that is, for a state x, (M_I·x)_i = Σ_j x_j + d_i·x_i, where d is the diagonal -// of the permutation (diag16 for width 16, diag24 for width 24). -func (p *Parameters) InternalMatrix() [][]fr.Element { - switch p.Width { - case 16: - return internalposeidon2.InternalMatrix(diag16) - case 24: - return internalposeidon2.InternalMatrix(diag24) - default: - panic("only Width=16,24 are supported") - } -} - // NewPermutation returns a new Poseidon2 permutation instance. func NewPermutation(t, rf, rp int) *Permutation { if t != 16 && t != 24 { diff --git a/field/koalabear/poseidon2/poseidon2_test.go b/field/koalabear/poseidon2/poseidon2_test.go index f0363fe834..c47680863f 100644 --- a/field/koalabear/poseidon2/poseidon2_test.go +++ b/field/koalabear/poseidon2/poseidon2_test.go @@ -64,52 +64,6 @@ func TestMulMulInternalInPlaceWidth24(t *testing.T) { } } -// denseMatrixMul returns m·x. -func denseMatrixMul(m [][]fr.Element, x []fr.Element) []fr.Element { - res := make([]fr.Element, len(m)) - for i := range m { - var tmp fr.Element - for j := range m[i] { - tmp.Mul(&m[i][j], &x[j]) - res[i].Add(&res[i], &tmp) - } - } - return res -} - -// TestDenseMatricesMatchInPlace checks that the dense matrices act on a random -// state as the in-place multiplications of the permutation. -func TestDenseMatricesMatchInPlace(t *testing.T) { - for _, tc := range []struct { - width, nbFullRounds, nbPartialRounds int - }{ - {16, 6, 21}, - {24, 6, 21}, - } { - h := NewPermutation(tc.width, tc.nbFullRounds, tc.nbPartialRounds) - - for name, step := range map[string]struct { - dense [][]fr.Element - inPlace func([]fr.Element) - }{ - "external": {h.params.ExternalMatrix(), h.matMulExternalInPlace}, - "internal": {h.params.InternalMatrix(), h.matMulInternalInPlace}, - } { - x := make([]fr.Element, tc.width) - for i := range x { - x[i].MustSetRandom() - } - want := denseMatrixMul(step.dense, x) - step.inPlace(x) - for i := range x { - if !x[i].Equal(&want[i]) { - t.Fatalf("width %d, %s matrix: dense and in-place multiplications differ at index %d", tc.width, name, i) - } - } - } - } -} - func TestAVX512Width16(t *testing.T) { if !cpu.SupportAVX512 { t.Skip("AVX512 not supported") diff --git a/field/mamabear/poseidon2/poseidon2.go b/field/mamabear/poseidon2/poseidon2.go index ecb25f7601..278bd39f60 100644 --- a/field/mamabear/poseidon2/poseidon2.go +++ b/field/mamabear/poseidon2/poseidon2.go @@ -15,7 +15,6 @@ import ( "golang.org/x/crypto/sha3" fr "github.com/consensys/gnark-crypto/field/mamabear" - internalposeidon2 "github.com/consensys/gnark-crypto/internal/poseidon2" ) var ( @@ -122,36 +121,6 @@ type Permutation struct { params *Parameters } -// ExternalMatrix returns the dense external matrix M_E of the permutation, which -// is applied to the state before the first round and after each full round. -// With I and J the identity and the all-ones matrices, and M4 the 4×4 block of -// Plonky3 (rows 2 3 1 1 / 1 2 3 1 / 1 1 2 3 / 3 1 1 2), M_E is -// -// - I + J for width 2 and 3; -// - M4 for width 4; -// - circ(2·M4, M4, ..., M4) = (I + J) ⊗ M4 for width 4k with k > 1. -func (p *Parameters) ExternalMatrix() [][]fr.Element { - return internalposeidon2.ExternalMatrix[fr.Element](p.Width, internalposeidon2.Plonky3) -} - -// InternalMatrix returns the dense internal matrix M_I of the permutation, which -// is applied to the state after each partial round. With J the all-ones matrix, -// -// M_I = J + diag(d), -// -// that is, for a state x, (M_I·x)_i = Σ_j x_j + d_i·x_i, where d is the diagonal -// of the permutation (diag16 for width 16, diag24 for width 24). -func (p *Parameters) InternalMatrix() [][]fr.Element { - switch p.Width { - case 16: - return internalposeidon2.InternalMatrix(diag16) - case 24: - return internalposeidon2.InternalMatrix(diag24) - default: - panic("only Width=16,24 are supported") - } -} - // NewPermutation returns a new Poseidon2 permutation instance. func NewPermutation(t, rf, rp int) *Permutation { if t != 16 && t != 24 { diff --git a/field/mamabear/poseidon2/poseidon2_test.go b/field/mamabear/poseidon2/poseidon2_test.go index fd59e1494c..f26d59464e 100644 --- a/field/mamabear/poseidon2/poseidon2_test.go +++ b/field/mamabear/poseidon2/poseidon2_test.go @@ -61,52 +61,6 @@ func TestMulMulInternalInPlaceWidth24(t *testing.T) { } } -// denseMatrixMul returns m·x. -func denseMatrixMul(m [][]fr.Element, x []fr.Element) []fr.Element { - res := make([]fr.Element, len(m)) - for i := range m { - var tmp fr.Element - for j := range m[i] { - tmp.Mul(&m[i][j], &x[j]) - res[i].Add(&res[i], &tmp) - } - } - return res -} - -// TestDenseMatricesMatchInPlace checks that the dense matrices act on a random -// state as the in-place multiplications of the permutation. -func TestDenseMatricesMatchInPlace(t *testing.T) { - for _, tc := range []struct { - width, nbFullRounds, nbPartialRounds int - }{ - {16, 8, 32}, - {24, 8, 32}, - } { - h := NewPermutation(tc.width, tc.nbFullRounds, tc.nbPartialRounds) - - for name, step := range map[string]struct { - dense [][]fr.Element - inPlace func([]fr.Element) - }{ - "external": {h.params.ExternalMatrix(), h.matMulExternalInPlace}, - "internal": {h.params.InternalMatrix(), h.matMulInternalInPlace}, - } { - x := make([]fr.Element, tc.width) - for i := range x { - x[i].MustSetRandom() - } - want := denseMatrixMul(step.dense, x) - step.inPlace(x) - for i := range x { - if !x[i].Equal(&want[i]) { - t.Fatalf("width %d, %s matrix: dense and in-place multiplications differ at index %d", tc.width, name, i) - } - } - } - } -} - func TestPoseidon2Width16(t *testing.T) { var input, expected [16]fr.Element // these are random values generated by MustSetRandom() diff --git a/internal/generator/crypto/hash/poseidon2/template/poseidon2.go.tmpl b/internal/generator/crypto/hash/poseidon2/template/poseidon2.go.tmpl index 0bab8fd649..e7c0ff37f3 100644 --- a/internal/generator/crypto/hash/poseidon2/template/poseidon2.go.tmpl +++ b/internal/generator/crypto/hash/poseidon2/template/poseidon2.go.tmpl @@ -5,7 +5,6 @@ import ( "golang.org/x/crypto/sha3" "github.com/consensys/gnark-crypto/ecc/{{ .Name }}/fr" - internalposeidon2 "github.com/consensys/gnark-crypto/internal/poseidon2" ) var ( @@ -148,50 +147,6 @@ type Permutation struct { params *Parameters } -// ExternalMatrix returns the dense external matrix M_E of the permutation, which -// is applied to the state before the first round and after each full round. -// With I and J the identity and the all-ones matrices, and M4 the 4×4 block of -// the Poseidon2 paper (rows 5 7 1 3 / 4 6 1 1 / 1 3 5 7 / 1 1 4 6), M_E is -// -// - I + J for width 2 and 3; -// - M4 for width 4; -// - circ(2·M4, M4, ..., M4) = (I + J) ⊗ M4 for width 4k with k > 1. -func (p *Parameters) ExternalMatrix() [][]fr.Element { - return internalposeidon2.ExternalMatrix[fr.Element](p.Width, internalposeidon2.Paper) -} - -// InternalMatrix returns the dense internal matrix M_I of the permutation, which -// is applied to the state after each partial round. With J the all-ones matrix, -// -// M_I = J + diag(d), -// -// that is, for a state x, (M_I·x)_i = Σ_j x_j + d_i·x_i, where d is the diagonal -// of the permutation{{- if eq .Name "bn254" }} (DiagM1 for the widths of 4 and more){{ end -}}. -func (p *Parameters) InternalMatrix() [][]fr.Element { - var d []fr.Element - switch p.Width { - case 2: - d = make([]fr.Element, 2) - d[0].SetOne() - d[1].SetUint64(2) - case 3: - d = make([]fr.Element, 3) - d[0].SetOne() - d[1].SetOne() - d[2].SetUint64(2) - default: - {{- if eq .Name "bn254" }} - if len(p.DiagM1) != p.Width { - panic("missing internal matrix diagonal") - } - d = p.DiagM1 - {{- else }} - panic("only Width=2,3 are supported") - {{- end }} - } - return internalposeidon2.InternalMatrix(d) -} - // NewPermutation returns a new Poseidon2 permutation instance. func NewPermutation(t, rf, rp int) *Permutation { {{- if eq .Name "bn254" }} diff --git a/internal/generator/crypto/hash/poseidon2/template/poseidon2.test.go.tmpl b/internal/generator/crypto/hash/poseidon2/template/poseidon2.test.go.tmpl index 46675afd29..115c22de9d 100644 --- a/internal/generator/crypto/hash/poseidon2/template/poseidon2.test.go.tmpl +++ b/internal/generator/crypto/hash/poseidon2/template/poseidon2.test.go.tmpl @@ -232,56 +232,4 @@ func TestHashReset(t *testing.T) { require.NoError(t, err) require.Equal(t, res, h.Sum(nil)) -} - -// denseMatrixMul returns m·x. -func denseMatrixMul(m [][]fr.Element, x []fr.Element) []fr.Element { - res := make([]fr.Element, len(m)) - for i := range m { - var tmp fr.Element - for j := range m[i] { - tmp.Mul(&m[i][j], &x[j]) - res[i].Add(&res[i], &tmp) - } - } - return res -} - -// TestDenseMatricesMatchInPlace checks that the dense matrices act on a random -// state as the in-place multiplications of the permutation. -func TestDenseMatricesMatchInPlace(t *testing.T) { - for _, tc := range []struct { - width, nbFullRounds, nbPartialRounds int - }{ - {2, 8, 56}, - {3, 8, 56}, -{{- if eq .Name "bn254" }} - {4, 8, 56}, - {8, 8, 57}, - {12, 8, 57}, - {16, 8, 57}, -{{- end }} - } { - h := NewPermutation(tc.width, tc.nbFullRounds, tc.nbPartialRounds) - - for name, step := range map[string]struct { - dense [][]fr.Element - inPlace func([]fr.Element) - }{ - "external": {h.params.ExternalMatrix(), h.matMulExternalInPlace}, - "internal": {h.params.InternalMatrix(), h.matMulInternalInPlace}, - } { - x := make([]fr.Element, tc.width) - for i := range x { - x[i].MustSetRandom() - } - want := denseMatrixMul(step.dense, x) - step.inPlace(x) - for i := range x { - if !x[i].Equal(&want[i]) { - t.Fatalf("width %d, %s matrix: dense and in-place multiplications differ at index %d", tc.width, name, i) - } - } - } - } -} +} \ No newline at end of file diff --git a/internal/generator/field/template/poseidon2/poseidon2.go.tmpl b/internal/generator/field/template/poseidon2/poseidon2.go.tmpl index f41eaacee2..7c3c2aef1c 100644 --- a/internal/generator/field/template/poseidon2/poseidon2.go.tmpl +++ b/internal/generator/field/template/poseidon2/poseidon2.go.tmpl @@ -9,7 +9,6 @@ import ( "golang.org/x/crypto/sha3" fr "{{ .FieldPackagePath }}" - internalposeidon2 "github.com/consensys/gnark-crypto/internal/poseidon2" {{- if .F31}} "github.com/consensys/gnark-crypto/utils/cpu" @@ -146,36 +145,6 @@ type Permutation struct { params *Parameters } -// ExternalMatrix returns the dense external matrix M_E of the permutation, which -// is applied to the state before the first round and after each full round. -// With I and J the identity and the all-ones matrices, and M4 the 4×4 block of -// {{if eq .FF "goldilocks"}}the Poseidon2 paper (rows 5 7 1 3 / 4 6 1 1 / 1 3 5 7 / 1 1 4 6){{else}}Plonky3 (rows 2 3 1 1 / 1 2 3 1 / 1 1 2 3 / 3 1 1 2){{end}}, M_E is -// -// - I + J for width 2 and 3; -// - M4 for width 4; -// - circ(2·M4, M4, ..., M4) = (I + J) ⊗ M4 for width 4k with k > 1. -func (p *Parameters) ExternalMatrix() [][]fr.Element { - return internalposeidon2.ExternalMatrix[fr.Element](p.Width, internalposeidon2.{{if eq .FF "goldilocks"}}Paper{{else}}Plonky3{{end}}) -} - -// InternalMatrix returns the dense internal matrix M_I of the permutation, which -// is applied to the state after each partial round. With J the all-ones matrix, -// -// M_I = J + diag(d), -// -// that is, for a state x, (M_I·x)_i = Σ_j x_j + d_i·x_i, where d is the diagonal -// of the permutation (diag{{$wc}} for width {{$wc}}, diag{{$ws}} for width {{$ws}}). -func (p *Parameters) InternalMatrix() [][]fr.Element { - switch p.Width { - case {{$wc}}: - return internalposeidon2.InternalMatrix(diag{{$wc}}) - case {{$ws}}: - return internalposeidon2.InternalMatrix(diag{{$ws}}) - default: - panic("only Width={{$wc}},{{$ws}} are supported") - } -} - // NewPermutation returns a new Poseidon2 permutation instance. func NewPermutation(t, rf, rp int) *Permutation { {{- if or (eq .Family "babybear") (eq .Family "koalabear")}} diff --git a/internal/generator/field/template/poseidon2/poseidon2_test.go.tmpl b/internal/generator/field/template/poseidon2/poseidon2_test.go.tmpl index 5fe2a4f1f5..8000db0248 100644 --- a/internal/generator/field/template/poseidon2/poseidon2_test.go.tmpl +++ b/internal/generator/field/template/poseidon2/poseidon2_test.go.tmpl @@ -68,53 +68,6 @@ func TestMulMulInternalInPlaceWidth{{- $w1}}(t *testing.T) { } -// denseMatrixMul returns m·x. -func denseMatrixMul(m [][]fr.Element, x []fr.Element) []fr.Element { - res := make([]fr.Element, len(m)) - for i := range m { - var tmp fr.Element - for j := range m[i] { - tmp.Mul(&m[i][j], &x[j]) - res[i].Add(&res[i], &tmp) - } - } - return res -} - -// TestDenseMatricesMatchInPlace checks that the dense matrices act on a random -// state as the in-place multiplications of the permutation. -func TestDenseMatricesMatchInPlace(t *testing.T) { - for _, tc := range []struct { - width, nbFullRounds, nbPartialRounds int - }{ - { {{- $w0}}, {{.ParamsCompression.FullRounds}}, {{.ParamsCompression.PartialRounds}} }, - { {{- $w1}}, {{.ParamsSponge.FullRounds}}, {{.ParamsSponge.PartialRounds}} }, - } { - h := NewPermutation(tc.width, tc.nbFullRounds, tc.nbPartialRounds) - - for name, step := range map[string]struct { - dense [][]fr.Element - inPlace func([]fr.Element) - }{ - "external": {h.params.ExternalMatrix(), h.matMulExternalInPlace}, - "internal": {h.params.InternalMatrix(), h.matMulInternalInPlace}, - } { - x := make([]fr.Element, tc.width) - for i := range x { - x[i].MustSetRandom() - } - want := denseMatrixMul(step.dense, x) - step.inPlace(x) - for i := range x { - if !x[i].Equal(&want[i]) { - t.Fatalf("width %d, %s matrix: dense and in-place multiplications differ at index %d", tc.width, name, i) - } - } - } - } -} - - {{- if .F31}} diff --git a/internal/poseidon2/matrices.go b/internal/poseidon2/matrices.go deleted file mode 100644 index 2129cf7fcc..0000000000 --- a/internal/poseidon2/matrices.go +++ /dev/null @@ -1,121 +0,0 @@ -// Copyright 2020-2026 Consensys Software Inc. -// Licensed under the Apache License, Version 2.0. See the LICENSE file for details. - -// Package poseidon2 holds the definitions of the linear layers of the Poseidon2 -// permutation that are shared by the generated field and curve packages. -package poseidon2 - -// ExternalMatrixKind names the source of the 4×4 block M4 from which the -// external matrix is built for the widths that are multiples of 4. -type ExternalMatrixKind int - -const ( - // Plonky3 is the block used by Plonky3, whose rows are - // - // 2 3 1 1 - // 1 2 3 1 - // 1 1 2 3 - // 3 1 1 2 - Plonky3 ExternalMatrixKind = iota - // Paper is the block of the Poseidon2 paper (https://eprint.iacr.org/2023/323.pdf, - // appendix B), whose rows are - // - // 5 7 1 3 - // 4 6 1 1 - // 1 3 5 7 - // 1 1 4 6 - Paper -) - -// m4Blocks holds the block M4 of each ExternalMatrixKind. -var m4Blocks = [...][4][4]int64{ - Plonky3: { - {2, 3, 1, 1}, - {1, 2, 3, 1}, - {1, 1, 2, 3}, - {3, 1, 1, 2}, - }, - Paper: { - {5, 7, 1, 3}, - {4, 6, 1, 1}, - {1, 3, 5, 7}, - {1, 1, 4, 6}, - }, -} - -// ring is the set of operations on a pointer PE to a ring element E that the -// matrix constructions need. -type ring[E any] interface { - *E - SetInt64(int64) *E - Add(*E, *E) *E -} - -// ExternalMatrix returns the dense width×width external matrix M_E, which the -// permutation applies before the first round and after each full round. With I -// and J the identity and the all-ones matrices, it is -// -// - I + J for width 2 and 3, whatever the kind; -// - M4 for width 4; -// - circ(2·M4, M4, ..., M4) = (I + J) ⊗ M4 for width 4k with k > 1, where -// I and J have size k. -// -// M4 is the block selected by kind. It panics for any other width, and for an -// unknown kind when a block is needed. -func ExternalMatrix[E any, PE ring[E]](width int, kind ExternalMatrixKind) [][]E { - m := newMatrix[E](width) - - switch { - case width == 2 || width == 3: - for i := range m { - for j := range m[i] { - v := int64(1) - if i == j { - v = 2 - } - PE(&m[i][j]).SetInt64(v) - } - } - return m - case width%4 != 0 || width <= 0: - panic("poseidon2: only widths 2, 3 and multiples of 4 are supported") - } - - if kind < 0 || int(kind) >= len(m4Blocks) { - panic("poseidon2: unknown external matrix kind") - } - m4 := &m4Blocks[kind] - for i := range m { - for j := range m[i] { - v := m4[i%4][j%4] - if width > 4 && i/4 == j/4 { - v *= 2 - } - PE(&m[i][j]).SetInt64(v) - } - } - return m -} - -// InternalMatrix returns the dense internal matrix M_I = J + diag(d) of width -// len(d), where J is the all-ones matrix: entry (i, j) is 1, plus d[i] when -// i == j. For a state x, (M_I·x)_i = Σ_j x_j + d[i]·x_i. The permutation applies -// M_I after each partial round. -func InternalMatrix[E any, PE ring[E]](d []E) [][]E { - m := newMatrix[E](len(d)) - for i := range m { - for j := range m[i] { - PE(&m[i][j]).SetInt64(1) - } - PE(&m[i][i]).Add(PE(&m[i][i]), PE(&d[i])) - } - return m -} - -func newMatrix[E any](n int) [][]E { - m := make([][]E, n) - for i := range m { - m[i] = make([]E, n) - } - return m -} diff --git a/internal/poseidon2/matrices_test.go b/internal/poseidon2/matrices_test.go deleted file mode 100644 index adb75dd840..0000000000 --- a/internal/poseidon2/matrices_test.go +++ /dev/null @@ -1,130 +0,0 @@ -// Copyright 2020-2026 Consensys Software Inc. -// Licensed under the Apache License, Version 2.0. See the LICENSE file for details. - -package poseidon2 - -import ( - "testing" - - fr "github.com/consensys/gnark-crypto/field/koalabear" -) - -func fromRows(rows [][]int64) [][]fr.Element { - m := newMatrix[fr.Element](len(rows)) - for i := range rows { - for j := range rows[i] { - m[i][j].SetInt64(rows[i][j]) - } - } - return m -} - -func requireEqual(t *testing.T, name string, got, want [][]fr.Element) { - t.Helper() - if len(got) != len(want) { - t.Fatalf("%s: got %d rows, want %d", name, len(got), len(want)) - } - for i := range want { - if len(got[i]) != len(want[i]) { - t.Fatalf("%s: row %d has %d entries, want %d", name, i, len(got[i]), len(want[i])) - } - for j := range want[i] { - if !got[i][j].Equal(&want[i][j]) { - t.Fatalf("%s: entry (%d, %d) is %s, want %s", name, i, j, got[i][j].String(), want[i][j].String()) - } - } - } -} - -func TestExternalMatrixSmallWidths(t *testing.T) { - for _, kind := range []ExternalMatrixKind{Plonky3, Paper} { - requireEqual(t, "width 2", ExternalMatrix[fr.Element](2, kind), fromRows([][]int64{ - {2, 1}, - {1, 2}, - })) - requireEqual(t, "width 3", ExternalMatrix[fr.Element](3, kind), fromRows([][]int64{ - {2, 1, 1}, - {1, 2, 1}, - {1, 1, 2}, - })) - } -} - -func TestExternalMatrixWidth4IsM4(t *testing.T) { - requireEqual(t, "plonky3", ExternalMatrix[fr.Element](4, Plonky3), fromRows([][]int64{ - {2, 3, 1, 1}, - {1, 2, 3, 1}, - {1, 1, 2, 3}, - {3, 1, 1, 2}, - })) - requireEqual(t, "paper", ExternalMatrix[fr.Element](4, Paper), fromRows([][]int64{ - {5, 7, 1, 3}, - {4, 6, 1, 1}, - {1, 3, 5, 7}, - {1, 1, 4, 6}, - })) -} - -// TestExternalMatrixCirculant checks circ(2·M4, M4, ..., M4) for the widths 4k, k > 1. -func TestExternalMatrixCirculant(t *testing.T) { - for _, kind := range []ExternalMatrixKind{Plonky3, Paper} { - for _, width := range []int{8, 12, 16, 24} { - got := ExternalMatrix[fr.Element](width, kind) - want := newMatrix[fr.Element](width) - for bi := range width / 4 { - for bj := range width / 4 { - for r := range 4 { - for c := range 4 { - v := m4Blocks[kind][r][c] - if bi == bj { - v *= 2 - } - want[4*bi+r][4*bj+c].SetInt64(v) - } - } - } - } - requireEqual(t, "circulant", got, want) - } - } -} - -func TestExternalMatrixPanics(t *testing.T) { - for _, width := range []int{-4, 0, 1, 5, 6, 7, 10} { - func() { - defer func() { - if recover() == nil { - t.Fatalf("width %d: expected a panic", width) - } - }() - ExternalMatrix[fr.Element](width, Plonky3) - }() - } - func() { - defer func() { - if recover() == nil { - t.Fatal("expected a panic for an unknown kind") - } - }() - ExternalMatrix[fr.Element](4, ExternalMatrixKind(7)) - }() -} - -func TestInternalMatrix(t *testing.T) { - d := make([]fr.Element, 4) - for i, v := range []int64{3, -1, 0, 7} { - d[i].SetInt64(v) - } - requireEqual(t, "internal", InternalMatrix(d), fromRows([][]int64{ - {4, 1, 1, 1}, - {1, 0, 1, 1}, - {1, 1, 1, 1}, - {1, 1, 1, 8}, - })) - // d is not modified - var want fr.Element - want.SetInt64(3) - if !d[0].Equal(&want) { - t.Fatal("InternalMatrix modified its argument") - } -}