diff --git a/libs/math/include/nil/crypto3/math/domains/arithmetic_sequence_domain.hpp b/libs/math/include/nil/crypto3/math/domains/arithmetic_sequence_domain.hpp index b7b909668..24f830e6a 100644 --- a/libs/math/include/nil/crypto3/math/domains/arithmetic_sequence_domain.hpp +++ b/libs/math/include/nil/crypto3/math/domains/arithmetic_sequence_domain.hpp @@ -56,21 +56,28 @@ namespace nil { std::vector arithmetic_sequence; field_value_type arithmetic_generator; + private: + // Queries derive domain points directly so they do not populate the transform precomputation. + field_value_type arithmetic_sequence_element(const std::size_t idx) const { + return arithmetic_generator * field_value_type(idx); + } + + public: void do_precomputation() { compute_subproduct_tree(this->subproduct_tree, log2(this->m)); - arithmetic_generator = field_value_type(fields::arithmetic_params::arithmetic_generator); - arithmetic_sequence = std::vector(this->m); for (std::size_t i = 0; i < arithmetic_sequence.size(); ++i) { - this->arithmetic_sequence[i] = this->arithmetic_generator * field_value_type(i); + this->arithmetic_sequence[i] = arithmetic_sequence_element(i); } precomputation_sentinel = true; } - arithmetic_sequence_domain(const std::size_t m) : evaluation_domain(m) { + arithmetic_sequence_domain(const std::size_t m) : + evaluation_domain(m), precomputation_sentinel(false), + arithmetic_generator(field_value_type(fields::arithmetic_params::arithmetic_generator)) { if (m <= 1) { throw std::invalid_argument("arithmetic(): expected m > 1"); } @@ -80,8 +87,6 @@ namespace nil { "arithmetic(): expected arithmetic_params::arithmetic_generator.is_zero() " "!= true"); } - - precomputation_sentinel = false; } void fft(std::vector &a) override { @@ -163,21 +168,18 @@ namespace nil { throw std::logic_error {"Not implemented yet"}; } - std::vector evaluate_all_lagrange_polynomials(const field_value_type &t) override { + std::vector + evaluate_all_lagrange_polynomials(const field_value_type &t) const override { /* Compute Lagrange polynomial of size m, with m+1 points (x_0, y_0), ... ,(x_m, y_m) */ /* Evaluate for x = t */ /* Return coeffs for each l_j(x) = (l / l_i[j]) * w[j] */ - if (!precomputation_sentinel) - do_precomputation(); - /** * If t equals one of the arithmetic progression values, * then output 1 at the right place, and 0 elsewhere. */ for (std::size_t i = 0; i < this->m; ++i) { - if (arithmetic_sequence[i] == t) // i.e., t equals this->arithmetic_sequence[i] - { + if (arithmetic_sequence_element(i) == t) { std::vector res(this->m, field_value_type::zero()); res[i] = field_value_type::one(); return res; @@ -189,15 +191,15 @@ namespace nil { * then compute each Lagrange coefficient. */ std::vector l(this->m); - l[0] = t - this->arithmetic_sequence[0]; + l[0] = t - arithmetic_sequence_element(0); field_value_type l_vanish = l[0]; field_value_type g_vanish = field_value_type::one(); for (std::size_t i = 1; i < this->m; i++) { - l[i] = t - this->arithmetic_sequence[i]; + l[i] = t - arithmetic_sequence_element(i); l_vanish *= l[i]; - g_vanish *= -this->arithmetic_sequence[i]; + g_vanish *= -arithmetic_sequence_element(i); } std::vector w(this->m); @@ -206,8 +208,8 @@ namespace nil { l[0] = l_vanish * l[0].inversed() * w[0]; for (std::size_t i = 1; i < this->m; i++) { field_value_type num = - this->arithmetic_sequence[i - 1] - this->arithmetic_sequence[this->m - 1]; - w[i] = w[i - 1] * num * this->arithmetic_sequence[i].inversed(); + arithmetic_sequence_element(i - 1) - arithmetic_sequence_element(this->m - 1); + w[i] = w[i - 1] * num * arithmetic_sequence_element(i).inversed(); l[i] = l_vanish * l[i].inversed() * w[i]; } @@ -216,7 +218,7 @@ namespace nil { std::vector evaluate_all_lagrange_polynomials( const typename std::vector::const_iterator &t_powers_begin, - const typename std::vector::const_iterator &t_powers_end) override { + const typename std::vector::const_iterator &t_powers_end) const override { if (std::size_t(std::distance(t_powers_begin, t_powers_end)) < this->m) { throw std::invalid_argument( "arithmetic_sequence_radix2: expected std::distance(t_powers_begin, t_powers_end) >= " @@ -227,16 +229,12 @@ namespace nil { /* Evaluate for x = t */ /* Return coeffs for each l_j(x) = (l / l_i[j]) * w[j] */ - if (!precomputation_sentinel) - do_precomputation(); - /** * If t equals one of the arithmetic progression values, * then output 1 at the right place, and 0 elsewhere. */ for (std::size_t i = 0; i < this->m; ++i) { - if (arithmetic_sequence[i] * t_powers_begin[0] == t_powers_begin[1]) // i.e., t equals a[i] - { + if (arithmetic_sequence_element(i) * t_powers_begin[0] == t_powers_begin[1]) { std::vector res(this->m, value_type::zero()); res[i] = t_powers_begin[0]; return res; @@ -248,16 +246,16 @@ namespace nil { * then compute each Lagrange coefficient. */ std::vector> l(this->m); - l[0] = polynomial({-arithmetic_sequence[0], field_value_type::one()}); + l[0] = polynomial({-arithmetic_sequence_element(0), field_value_type::one()}); ; polynomial l_vanish = l[0]; field_value_type g_vanish = field_value_type::one(); for (std::size_t i = 1; i < this->m; i++) { - l[i] = polynomial({-arithmetic_sequence[i], field_value_type::one()}); + l[i] = polynomial({-arithmetic_sequence_element(i), field_value_type::one()}); l_vanish = l_vanish * l[i]; - g_vanish *= -this->arithmetic_sequence[i]; + g_vanish *= -arithmetic_sequence_element(i); } std::vector w(this->m); @@ -276,8 +274,8 @@ namespace nil { for (std::size_t i = 1; i < this->m; i++) { field_value_type num = - this->arithmetic_sequence[i - 1] - this->arithmetic_sequence[this->m - 1]; - w[i] = w[i - 1] * num * this->arithmetic_sequence[i].inversed(); + arithmetic_sequence_element(i - 1) - arithmetic_sequence_element(this->m - 1); + w[i] = w[i - 1] * num * arithmetic_sequence_element(i).inversed(); for (std::size_t j = 0; j < l[i].size(); ++j) { result[i] = result[i] + t_powers_begin[j] * l[i][j]; @@ -289,36 +287,28 @@ namespace nil { } // This one is not the unity root actually, but it's ok for our purposes. - const field_value_type &get_unity_root() override { + const field_value_type &get_unity_root() const override { return arithmetic_generator; } - field_value_type get_domain_element(const std::size_t idx) override { - if (!this->precomputation_sentinel) - do_precomputation(); - - return this->arithmetic_sequence[idx]; + field_value_type get_domain_element(const std::size_t idx) const override { + return arithmetic_sequence_element(idx); } - field_value_type compute_vanishing_polynomial(const field_value_type &t) override { - if (!this->precomputation_sentinel) - do_precomputation(); - + field_value_type compute_vanishing_polynomial(const field_value_type &t) const override { /* Notes: Z = prod_{i = 0 to m} (t - a[i]) */ field_value_type Z = field_value_type::one(); for (std::size_t i = 0; i < this->m; i++) { - Z *= (t - this->arithmetic_sequence[i]); + Z *= (t - arithmetic_sequence_element(i)); } return Z; } - polynomial get_vanishing_polynomial() override { - if (!precomputation_sentinel) - do_precomputation(); - + polynomial get_vanishing_polynomial() const override { polynomial z({field_value_type::one()}); for (std::size_t i = 0; i < this->m; i++) { - z = z * polynomial({-arithmetic_sequence[i], field_value_type::one()}); + z = z * + polynomial({-arithmetic_sequence_element(i), field_value_type::one()}); } return z; } diff --git a/libs/math/include/nil/crypto3/math/domains/basic_radix2_domain.hpp b/libs/math/include/nil/crypto3/math/domains/basic_radix2_domain.hpp index 9d81600d1..ac85bba87 100644 --- a/libs/math/include/nil/crypto3/math/domains/basic_radix2_domain.hpp +++ b/libs/math/include/nil/crypto3/math/domains/basic_radix2_domain.hpp @@ -60,6 +60,23 @@ namespace nil { detail::create_fft_cache(this->m, omega.inversed(), fft_cache->second); } + void inverse_fft_impl(std::vector &a) const { + if (a.size() != this->m) { + if (a.size() < this->m) { + a.resize(this->m, value_type::zero()); + } else { + throw std::invalid_argument("basic_radix2: expected a.size() == this->m"); + } + } + + detail::basic_radix2_fft_cached(a, fft_cache->second); + + const field_value_type sconst = field_value_type(this->m).inversed(); + for (value_type &a_i : a) { + a_i *= sconst; + } + } + public: typedef FieldType field_type; using evaluation_domain::evaluate_all_lagrange_polynomials; @@ -148,51 +165,39 @@ namespace nil { } void inverse_fft(std::vector &a) override { - if (a.size() != this->m) { - if (a.size() < this->m) { - a.resize(this->m, value_type::zero()); - } else { - throw std::invalid_argument("basic_radix2: expected a.size() == this->m"); - } - } - - detail::basic_radix2_fft_cached(a, fft_cache->second); - - const field_value_type sconst = field_value_type(this->m).inversed(); - for (value_type &a_i : a) { - a_i *= sconst; - } + inverse_fft_impl(a); } - std::vector evaluate_all_lagrange_polynomials(const field_value_type &t) override { + std::vector + evaluate_all_lagrange_polynomials(const field_value_type &t) const override { return detail::basic_radix2_evaluate_all_lagrange_polynomials(this->m, t); } std::vector evaluate_all_lagrange_polynomials( const typename std::vector::const_iterator &t_powers_begin, - const typename std::vector::const_iterator &t_powers_end) override { + const typename std::vector::const_iterator &t_powers_end) const override { if (std::size_t(std::distance(t_powers_begin, t_powers_end)) < this->m) { throw std::invalid_argument( "basic_radix2: expected std::distance(t_powers_begin, t_powers_end) >= this->m"); } std::vector tmp(t_powers_begin, t_powers_begin + this->m); - this->inverse_fft(tmp); + inverse_fft_impl(tmp); return tmp; } - const field_value_type &get_unity_root() override { + const field_value_type &get_unity_root() const override { return omega; } - field_value_type get_domain_element(const std::size_t idx) override { + field_value_type get_domain_element(const std::size_t idx) const override { return omega.pow(idx); } - field_value_type compute_vanishing_polynomial(const field_value_type &t) override { + field_value_type compute_vanishing_polynomial(const field_value_type &t) const override { return (t.pow(this->m)) - field_value_type::one(); } - polynomial get_vanishing_polynomial() override { + polynomial get_vanishing_polynomial() const override { polynomial z(this->m + 1, field_value_type::zero()); z[this->m] = field_value_type::one(); z[0] = -field_value_type::one(); diff --git a/libs/math/include/nil/crypto3/math/domains/evaluation_domain.hpp b/libs/math/include/nil/crypto3/math/domains/evaluation_domain.hpp index edb71611f..ab12b9561 100644 --- a/libs/math/include/nil/crypto3/math/domains/evaluation_domain.hpp +++ b/libs/math/include/nil/crypto3/math/domains/evaluation_domain.hpp @@ -69,12 +69,12 @@ namespace nil { /** * Get the unity root. */ - virtual const field_value_type &get_unity_root() = 0; + virtual const field_value_type &get_unity_root() const = 0; /** * Get the idx-th element in S. */ - virtual field_value_type get_domain_element(const std::size_t idx) = 0; + virtual field_value_type get_domain_element(const std::size_t idx) const = 0; /** * Compute the FFT, over the domain S, of the vector a. @@ -105,7 +105,8 @@ namespace nil { * The output is a vector (b_{0},...,b_{m-1}) * where b_{i} is the evaluation of L_{i,S}(z) at z = t. */ - virtual std::vector evaluate_all_lagrange_polynomials(const field_value_type &t) = 0; + virtual std::vector + evaluate_all_lagrange_polynomials(const field_value_type &t) const = 0; /** * Evaluate all Lagrange polynomials and the domain vanishing polynomial at t. @@ -115,7 +116,7 @@ namespace nil { */ virtual std::vector evaluate_all_lagrange_polynomials(const field_value_type &t, - field_value_type &vanishing_polynomial_at_t) { + field_value_type &vanishing_polynomial_at_t) const { std::vector result = evaluate_all_lagrange_polynomials(t); vanishing_polynomial_at_t = compute_vanishing_polynomial(t); return result; @@ -132,17 +133,17 @@ namespace nil { */ virtual std::vector evaluate_all_lagrange_polynomials( const typename std::vector::const_iterator &t_powers_begin, - const typename std::vector::const_iterator &t_powers_end) = 0; + const typename std::vector::const_iterator &t_powers_end) const = 0; /** * Evaluate the vanishing polynomial of S at the field element t. */ - virtual field_value_type compute_vanishing_polynomial(const field_value_type &t) = 0; + virtual field_value_type compute_vanishing_polynomial(const field_value_type &t) const = 0; /** * Build the vanishing polynomial of S. */ - virtual polynomial get_vanishing_polynomial() = 0; + virtual polynomial get_vanishing_polynomial() const = 0; /** * Add the coefficients of the vanishing polynomial of S to the coefficients of the polynomial H. diff --git a/libs/math/include/nil/crypto3/math/domains/extended_radix2_domain.hpp b/libs/math/include/nil/crypto3/math/domains/extended_radix2_domain.hpp index 65508544e..1472c320c 100644 --- a/libs/math/include/nil/crypto3/math/domains/extended_radix2_domain.hpp +++ b/libs/math/include/nil/crypto3/math/domains/extended_radix2_domain.hpp @@ -157,7 +157,8 @@ namespace nil { throw std::logic_error {"Not implemented yet"}; } - std::vector evaluate_all_lagrange_polynomials(const field_value_type &t) override { + std::vector + evaluate_all_lagrange_polynomials(const field_value_type &t) const override { const std::vector T0 = detail::basic_radix2_evaluate_all_lagrange_polynomials(small_m, t); const std::vector T1 = @@ -181,7 +182,7 @@ namespace nil { std::vector evaluate_all_lagrange_polynomials( const typename std::vector::const_iterator &t_powers_begin, - const typename std::vector::const_iterator &t_powers_end) override { + const typename std::vector::const_iterator &t_powers_end) const override { if (std::size_t(std::distance(t_powers_begin, t_powers_end)) < this->m) { throw std::invalid_argument( "extended_radix2: expected std::distance(t_powers_begin, t_powers_end) >= this->m"); @@ -223,11 +224,11 @@ namespace nil { return result; } - const field_value_type &get_unity_root() override { + const field_value_type &get_unity_root() const override { return omega; } - field_value_type get_domain_element(const std::size_t idx) override { + field_value_type get_domain_element(const std::size_t idx) const override { if (idx < small_m) { return omega.pow(idx); } else { @@ -235,11 +236,11 @@ namespace nil { } } - field_value_type compute_vanishing_polynomial(const field_value_type &t) override { + field_value_type compute_vanishing_polynomial(const field_value_type &t) const override { return (t.pow(small_m) - field_value_type::one()) * (t.pow(small_m) - shift.pow(small_m)); } - polynomial get_vanishing_polynomial() override { + polynomial get_vanishing_polynomial() const override { polynomial z(2 * small_m + 1, field_value_type::zero()); field_value_type shift_to_small_m = shift.pow(small_m); z[2 * small_m] = field_value_type::one(); diff --git a/libs/math/include/nil/crypto3/math/domains/geometric_sequence_domain.hpp b/libs/math/include/nil/crypto3/math/domains/geometric_sequence_domain.hpp index 3d7fa673e..8dfbc1efc 100644 --- a/libs/math/include/nil/crypto3/math/domains/geometric_sequence_domain.hpp +++ b/libs/math/include/nil/crypto3/math/domains/geometric_sequence_domain.hpp @@ -300,7 +300,7 @@ namespace nil { */ std::vector evaluate_all_lagrange_polynomials(const field_value_type &t, - field_value_type &vanishing_polynomial_at_t) override { + field_value_type &vanishing_polynomial_at_t) const override { std::vector denominators(this->m, field_value_type::zero()); vanishing_polynomial_at_t = field_value_type::one(); for (std::size_t i = 0; i < this->m; ++i) { @@ -324,14 +324,15 @@ namespace nil { return result; } - std::vector evaluate_all_lagrange_polynomials(const field_value_type &t) override { + std::vector + evaluate_all_lagrange_polynomials(const field_value_type &t) const override { field_value_type vanishing_polynomial_at_t; return evaluate_all_lagrange_polynomials(t, vanishing_polynomial_at_t); } std::vector evaluate_all_lagrange_polynomials( const typename std::vector::const_iterator &t_powers_begin, - const typename std::vector::const_iterator &t_powers_end) override { + const typename std::vector::const_iterator &t_powers_end) const override { if (std::size_t(std::distance(t_powers_begin, t_powers_end)) < this->m) { throw std::invalid_argument( "geometric_sequence_radix2: expected std::distance(t_powers_begin, t_powers_end) >= " @@ -417,15 +418,15 @@ namespace nil { // The base interface requires this legacy name. Geometric domains return r, which need not be a root // of unity. - const field_value_type &get_unity_root() override { + const field_value_type &get_unity_root() const override { return precomputation_.generator; } - field_value_type get_domain_element(const std::size_t idx) override { + field_value_type get_domain_element(const std::size_t idx) const override { return precomputation_.geometric_sequence[idx]; } - field_value_type compute_vanishing_polynomial(const field_value_type &t) override { + field_value_type compute_vanishing_polynomial(const field_value_type &t) const override { // Evaluate the domain vanishing polynomial Z(t) = product_(i=0)^(m-1) (t - x_i). field_value_type Z = field_value_type::one(); for (std::size_t i = 0; i < this->m; i++) { @@ -434,7 +435,7 @@ namespace nil { return Z; } - polynomial get_vanishing_polynomial() override { + polynomial get_vanishing_polynomial() const override { return precomputation_.vanishing_polynomial; } diff --git a/libs/math/include/nil/crypto3/math/domains/step_radix2_domain.hpp b/libs/math/include/nil/crypto3/math/domains/step_radix2_domain.hpp index b1c61d2c7..298d142bb 100644 --- a/libs/math/include/nil/crypto3/math/domains/step_radix2_domain.hpp +++ b/libs/math/include/nil/crypto3/math/domains/step_radix2_domain.hpp @@ -197,7 +197,8 @@ namespace nil { throw std::logic_error {"Not implemented yet"}; } - std::vector evaluate_all_lagrange_polynomials(const field_value_type &t) override { + std::vector + evaluate_all_lagrange_polynomials(const field_value_type &t) const override { std::vector inner_big = detail::basic_radix2_evaluate_all_lagrange_polynomials(big_m, t); std::vector inner_small = @@ -227,7 +228,7 @@ namespace nil { std::vector evaluate_all_lagrange_polynomials( const typename std::vector::const_iterator &t_powers_begin, - const typename std::vector::const_iterator &t_powers_end) override { + const typename std::vector::const_iterator &t_powers_end) const override { if (std::size_t(std::distance(t_powers_begin, t_powers_end)) < this->m) { throw std::invalid_argument( "extended_radix2: expected std::distance(t_powers_begin, t_powers_end) >= this->m"); @@ -277,11 +278,11 @@ namespace nil { return result; } - const field_value_type &get_unity_root() override { + const field_value_type &get_unity_root() const override { return omega; } - field_value_type get_domain_element(const std::size_t idx) override { + field_value_type get_domain_element(const std::size_t idx) const override { if (idx < big_m) { return big_omega.pow(idx); } else { @@ -289,11 +290,11 @@ namespace nil { } } - field_value_type compute_vanishing_polynomial(const field_value_type &t) override { + field_value_type compute_vanishing_polynomial(const field_value_type &t) const override { return (t.pow(big_m) - field_value_type::one()) * (t.pow(small_m) - omega.pow(small_m)); } - polynomial get_vanishing_polynomial() override { + polynomial get_vanishing_polynomial() const override { polynomial z(big_m + small_m + 1, field_value_type::zero()); field_value_type omega_to_small_m = omega.pow(small_m); z[big_m + small_m] = field_value_type::one(); diff --git a/libs/math/test/geometric_sequence_domain.cpp b/libs/math/test/geometric_sequence_domain.cpp index 5de6de0b3..8bd080271 100644 --- a/libs/math/test/geometric_sequence_domain.cpp +++ b/libs/math/test/geometric_sequence_domain.cpp @@ -159,6 +159,34 @@ BOOST_AUTO_TEST_CASE(combined_lagrange_api_has_a_default_domain_implementation) BOOST_CHECK_EQUAL(vanishing_at_t, domain.compute_vanishing_polynomial(t)); } +BOOST_AUTO_TEST_CASE(query_operations_work_through_a_const_geometric_domain) { + using value_type = bn254_fq::value_type; + + constexpr std::size_t domain_size = 9; + math::geometric_sequence_domain reference_domain(domain_size); + const math::geometric_sequence_domain domain(domain_size); + const math::evaluation_domain &abstract_domain = domain; + + BOOST_CHECK_EQUAL(abstract_domain.get_unity_root(), reference_domain.get_unity_root()); + + const std::size_t point_index = 4; + BOOST_CHECK_EQUAL(abstract_domain.get_domain_element(point_index), + reference_domain.get_domain_element(point_index)); + + const value_type t = value_type(19u); + value_type vanishing_at_t; + value_type reference_vanishing_at_t; + const std::vector weights = abstract_domain.evaluate_all_lagrange_polynomials(t, vanishing_at_t); + const std::vector reference_weights = + reference_domain.evaluate_all_lagrange_polynomials(t, reference_vanishing_at_t); + BOOST_CHECK(weights == reference_weights); + BOOST_CHECK_EQUAL(vanishing_at_t, reference_vanishing_at_t); + + BOOST_CHECK_EQUAL(abstract_domain.compute_vanishing_polynomial(t), + reference_domain.compute_vanishing_polynomial(t)); + BOOST_CHECK(abstract_domain.get_vanishing_polynomial() == reference_domain.get_vanishing_polynomial()); +} + BOOST_AUTO_TEST_CASE(lagrange_weights_match_the_definition) { using value_type = bn254_fq::value_type;