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>
This commit is contained in:
parent
695f8e84f7
commit
0d22946a0e
1835 changed files with 170291 additions and 2849 deletions
277
include/grotto/dwt_lut.hpp
Normal file
277
include/grotto/dwt_lut.hpp
Normal file
|
|
@ -0,0 +1,277 @@
|
|||
/// @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__
|
||||
Loading…
Add table
Add a link
Reference in a new issue