/// @file grotto/dwt_lut.hpp /// @brief Haar and bior(5,3) compressed lookup tables. /// @details Cleartext evaluators for Reis, Ugurbil, Wagh, Henry, and de Vega, /// PoPETs 2025 ([ePrint 2025/013](@ref bib_wave)). A power-of-two /// grid is reduced by a discrete wavelet transform and stored as /// fixed-point approximation coefficients. Haar evaluates with one /// lookup of the high bits. bior(5,3) evaluates Equation (8): two /// adjacent coefficients, weighted by `(2^j - lsb)` and `lsb`. /// @copyright Copyright (c) 2019-2026 Ryan Henry and [others](@ref authors) /// @license Released under a GNU General Public v2.0 (GPLv2) license. #ifndef LIBDPF_INCLUDE_GROTTO_DWT_LUT_HPP__ #define LIBDPF_INCLUDE_GROTTO_DWT_LUT_HPP__ #include #include #include #include #include namespace grotto { /// @brief Which wavelet compresses the table. enum class dwt_family : unsigned { /// Orthogonal Haar. One approximation coefficient per block of `2^depth` samples. haar = 0, /// Biorthogonal bior(5,3), PyWavelets `bior2.2`. Two taps at evaluation. bior53, }; /// @brief Fixed-point Haar or bior(5,3) table. /// @details `coeff` holds the quantized approximation coefficients. /// `operator()` is the cleartext evaluation those MPC protocols return. struct dwt_lut { dwt_family family = dwt_family::haar; /// Transform depth `j`. The kept index is `raw >> depth`. unsigned depth = 0; /// `log2` of the sample count. Inputs lie in `[0, 2^{domain_bits})`. unsigned domain_bits = 0; /// Fraction bits of every stored coefficient and of the returned word. unsigned fractional_bits = 0; std::vector coeff; /// \complexity Haar is one indexing step, `Θ(1)`. bior(5,3) is two multiplications and a shift by `2j`, `Θ(1)`. Extra space `Θ(1)`. /// @param raw index on the sampled grid, in `[0, 2^{domain_bits})` /// @return fixed-point approximation, `fractional_bits` fraction bits std::int64_t operator()(std::uint64_t raw) const { if (domain_bits == 0 || domain_bits > 30 || depth == 0 || depth > domain_bits) throw std::logic_error("dwt lut: table is not initialized"); const std::uint64_t domain = std::uint64_t{1} << domain_bits; if (raw >= domain) throw std::out_of_range("dwt lut: input is outside the sampled domain"); const std::uint64_t msb = raw >> depth; if (family == dwt_family::haar) { if (msb >= coeff.size()) throw std::out_of_range("dwt lut: Haar index"); return coeff[static_cast(msb)]; } if (coeff.empty() || msb >= coeff.size()) throw std::out_of_range("dwt lut: bior index"); // Smooth extension prepends two coefficients, so bin `msb` lives at // `msb + 2` and its neighbor at `msb + 3` (Equation (8), artifact indexing). const std::size_t n = coeff.size(); const auto c0 = coeff[static_cast((msb + 2) % n)]; const auto c1 = coeff[static_cast((msb + 3) % n)]; const std::uint64_t span = std::uint64_t{1} << depth; const std::uint64_t lsb = raw & (span - 1); const __int128 comb = __int128(c0) * static_cast<__int128>(span - lsb) + __int128(c1) * static_cast<__int128>(lsb); const __int128 den = __int128{1} << (2 * depth); return detail_floor_div(comb, den); } private: static std::int64_t detail_floor_div(__int128 num, __int128 den) { __int128 q = num / den; const __int128 r = num % den; if (r != 0 && num < 0) --q; if (q > std::numeric_limits::max() || q < std::numeric_limits::min()) throw std::overflow_error("dwt lut: value does not fit int64"); return static_cast(q); } }; namespace dwt_detail { inline unsigned log2_pow2(std::size_t n) { if (n < 2 || (n & (n - 1)) != 0) throw std::invalid_argument("dwt lut: sample count must be a power of two, at least 2"); unsigned bits = 0; while ((std::size_t{1} << bits) != n) ++bits; if (bits > 30) throw std::invalid_argument("dwt lut: sample count exceeds 2^30"); return bits; } inline std::int64_t quantize(double y, unsigned fractional_bits) { const double scaled = std::floor(y * std::ldexp(1.0, static_cast(fractional_bits))); const double lo = static_cast(std::numeric_limits::min()); if (!(scaled >= lo && scaled < 0x1p63)) throw std::overflow_error("dwt lut: coefficient does not fit int64"); return static_cast(scaled); } inline double half_exp(int k) { double scale = std::ldexp(1.0, k / 2); if (k % 2 != 0) scale *= (k < 0) ? std::sqrt(0.5) : std::sqrt(2.0); return scale; } inline std::vector down_approx(const std::vector & in, const double * filt, int F) { const int N = static_cast(in.size()); if (N < 1) throw std::invalid_argument("dwt lut: empty approximation"); const int expect = (N + F - 1) / 2; if (N < 2) { double sum = 0; for (int j = 0; j < F; ++j) sum += filt[j]; return std::vector(static_cast(expect), in[0] * sum); } std::vector out; out.reserve(static_cast(expect)); int i = 1; while (i < F && i < N) { double s = 0; int j = 0; for (; j <= i; ++j) s += filt[j] * in[static_cast(i - j)]; for (int k = 1; j < F; ++j, ++k) s += filt[j] * (in[0] + static_cast(k) * (in[0] - in[1])); out.push_back(s); i += 2; } while (i < N) { double s = 0; for (int j = 0; j < F; ++j) s += in[static_cast(i - j)] * filt[j]; out.push_back(s); i += 2; } while (i < F) { double s = 0; int j = 0; for (int k = i - N + 1; i - j >= N; ++j, --k) s += filt[j] * (in[static_cast(N - 1)] + static_cast(k) * (in[static_cast(N - 1)] - in[static_cast(N - 2)])); for (; j <= i; ++j) s += filt[j] * in[static_cast(i - j)]; for (int k = 1; j < F; ++j, ++k) s += filt[j] * (in[0] + static_cast(k) * (in[0] - in[1])); out.push_back(s); i += 2; } while (i < N + F - 1) { double s = 0; int j = 0; for (int k = i - N + 1; i - j >= N; ++j, --k) s += filt[j] * (in[static_cast(N - 1)] + static_cast(k) * (in[static_cast(N - 1)] - in[static_cast(N - 2)])); for (; j < F; ++j) s += filt[j] * in[static_cast(i - j)]; out.push_back(s); i += 2; } if (static_cast(out.size()) != expect) throw std::logic_error("dwt lut: approximation length"); return out; } inline dwt_lut build(dwt_family family, const std::vector & samples, unsigned fractional_bits, unsigned depth) { if (fractional_bits > 62) throw std::invalid_argument("dwt lut: fractional width does not fit"); const unsigned domain_bits = log2_pow2(samples.size()); if (depth == 0 || depth > domain_bits) throw std::invalid_argument("dwt lut: depth must lie in 1 .. domain bits"); const double s2 = std::sqrt(2.0); const double haar_lo[2] = {s2 / 2, s2 / 2}; const double s2_8 = s2 / 8; const double s2_4 = s2 / 4; const double bior_lo[6] = {0.0, -s2_8, s2_4, 3 * s2_4, s2_4, -s2_8}; const double * filt = haar_lo; int flen = 2; int scale_sign = -1; if (family == dwt_family::bior53) { filt = bior_lo; flen = 6; scale_sign = 1; } std::vector approx = samples; for (unsigned level = 0; level < depth; ++level) approx = down_approx(approx, filt, flen); const double scale = half_exp(scale_sign * static_cast(depth)); dwt_lut table; table.family = family; table.depth = depth; table.domain_bits = domain_bits; table.fractional_bits = fractional_bits; table.coeff.reserve(approx.size()); for (double a : approx) table.coeff.push_back(quantize(a * scale, fractional_bits)); return table; } } // namespace dwt_detail /// @brief Sample `f` on `i · 2^{-fractional_bits}` for `i` in `[0, 2^{domain_bits})`. /// \complexity `Θ(2^{domain_bits})` calls and extra space. /// @tparam Fn callable `double(double)` template std::vector sample_dwt_signal(unsigned domain_bits, unsigned fractional_bits, Fn && f) { if (domain_bits == 0 || domain_bits > 30) throw std::invalid_argument("dwt lut: domain bits must lie in 1 .. 30"); if (fractional_bits > 62) throw std::invalid_argument("dwt lut: fractional width does not fit"); const std::size_t n = std::size_t{1} << domain_bits; const double step = std::ldexp(1.0, -static_cast(fractional_bits)); std::vector samples(n); for (std::size_t i = 0; i < n; ++i) samples[i] = f(static_cast(i) * step); return samples; } /// \complexity One smooth-extension low-pass per level. `Θ(N)` arithmetic and extra space, `N` the sample count. /// @param samples real signal, length `2^{domain bits}` /// @param fractional_bits fraction bits of the stored words /// @param depth transform depth `j`, in `1 .. log2(samples)` dwt_lut make_haar_dwt_lut(const std::vector & samples, unsigned fractional_bits, unsigned depth) { return dwt_detail::build(dwt_family::haar, samples, fractional_bits, depth); } /// \complexity One smooth-extension low-pass per level. `Θ(N)` arithmetic and extra space, `N` the sample count. /// @param samples real signal, length `2^{domain bits}` /// @param fractional_bits fraction bits of the stored words /// @param depth transform depth `j`, in `1 .. log2(samples)` dwt_lut make_bior53_dwt_lut(const std::vector & samples, unsigned fractional_bits, unsigned depth) { return dwt_detail::build(dwt_family::bior53, samples, fractional_bits, depth); } } // namespace grotto #endif // LIBDPF_INCLUDE_GROTTO_DWT_LUT_HPP__