Ship the TLS mesh, composer, Beaver/Yao/leaf MPC, prep/online paths, apps, and docs so the tree is pushable before elevating share_expr, security_mode, and prep resume. Co-authored-by: Cursor <cursoragent@cursor.com>
277 lines
10 KiB
C++
277 lines
10 KiB
C++
/// @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 <cmath>
|
|
#include <cstdint>
|
|
#include <limits>
|
|
#include <stdexcept>
|
|
#include <vector>
|
|
|
|
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<std::int64_t> 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<std::size_t>(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<std::size_t>((msb + 2) % n)];
|
|
const auto c1 = coeff[static_cast<std::size_t>((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<std::int64_t>::max()
|
|
|| q < std::numeric_limits<std::int64_t>::min())
|
|
throw std::overflow_error("dwt lut: value does not fit int64");
|
|
return static_cast<std::int64_t>(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<int>(fractional_bits)));
|
|
const double lo = static_cast<double>(std::numeric_limits<std::int64_t>::min());
|
|
if (!(scaled >= lo && scaled < 0x1p63))
|
|
throw std::overflow_error("dwt lut: coefficient does not fit int64");
|
|
return static_cast<std::int64_t>(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<double> down_approx(const std::vector<double> & in, const double * filt, int F)
|
|
{
|
|
const int N = static_cast<int>(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<double>(static_cast<std::size_t>(expect), in[0] * sum);
|
|
}
|
|
|
|
std::vector<double> out;
|
|
out.reserve(static_cast<std::size_t>(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<std::size_t>(i - j)];
|
|
for (int k = 1; j < F; ++j, ++k)
|
|
s += filt[j] * (in[0] + static_cast<double>(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<std::size_t>(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<std::size_t>(N - 1)]
|
|
+ static_cast<double>(k) * (in[static_cast<std::size_t>(N - 1)]
|
|
- in[static_cast<std::size_t>(N - 2)]));
|
|
for (; j <= i; ++j)
|
|
s += filt[j] * in[static_cast<std::size_t>(i - j)];
|
|
for (int k = 1; j < F; ++j, ++k)
|
|
s += filt[j] * (in[0] + static_cast<double>(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<std::size_t>(N - 1)]
|
|
+ static_cast<double>(k) * (in[static_cast<std::size_t>(N - 1)]
|
|
- in[static_cast<std::size_t>(N - 2)]));
|
|
for (; j < F; ++j)
|
|
s += filt[j] * in[static_cast<std::size_t>(i - j)];
|
|
out.push_back(s);
|
|
i += 2;
|
|
}
|
|
if (static_cast<int>(out.size()) != expect)
|
|
throw std::logic_error("dwt lut: approximation length");
|
|
return out;
|
|
}
|
|
|
|
inline dwt_lut build(dwt_family family, const std::vector<double> & 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<double> approx = samples;
|
|
for (unsigned level = 0; level < depth; ++level)
|
|
approx = down_approx(approx, filt, flen);
|
|
|
|
const double scale = half_exp(scale_sign * static_cast<int>(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 <typename Fn>
|
|
std::vector<double> 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<int>(fractional_bits));
|
|
std::vector<double> samples(n);
|
|
for (std::size_t i = 0; i < n; ++i)
|
|
samples[i] = f(static_cast<double>(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<double> & 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<double> & 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__
|