libdpf/include/grotto/dwt_lut.hpp
Ryan Henry 0d22946a0e Checkpoint the party/runtime stack before share-program and malicious-mode work.
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>
2026-09-28 05:59:19 -06:00

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__