#include "SPECK_FLT.h"
#include <algorithm>
#include <cassert>
#include <cfenv>
#include <cfloat>
#include <cmath>
#include <cstring>
#include <numeric>
template <typename T>
void sperr::SPECK_FLT::copy_data(const T* p, size_t len)
{
static_assert(std::is_floating_point<T>::value, "!! Only floating point values are supported !!");
m_vals_d.resize(len);
std::copy(p, p + len, m_vals_d.begin());
}
template void sperr::SPECK_FLT::copy_data(const double*, size_t);
template void sperr::SPECK_FLT::copy_data(const float*, size_t);
void sperr::SPECK_FLT::take_data(sperr::vecd_type&& buf)
{
m_vals_d = std::move(buf);
}
auto sperr::SPECK_FLT::use_bitstream(const void* p, size_t len) -> RTNType
{
m_vals_d.clear();
m_sign_array.resize(0);
std::visit([](auto&& vec) { vec.clear(); }, m_vals_ui);
m_q = 0.0;
m_has_outlier = false;
const auto* const ptr = static_cast<const uint8_t*>(p);
if (len < m_condi_bitstream.size())
return RTNType::WrongLength;
std::copy(ptr, ptr + m_condi_bitstream.size(), m_condi_bitstream.begin());
if (m_conditioner.is_constant(m_condi_bitstream[0])) {
if (len == m_condi_bitstream.size())
return RTNType::Good;
else
return RTNType::WrongLength;
}
else {
m_q = m_conditioner.retrieve_q(m_condi_bitstream);
assert(m_q > 0.0);
}
auto pos = m_condi_bitstream.size();
auto remaining_len = len - pos;
assert(remaining_len >= SPECK_INT<uint8_t>::header_size);
const uint8_t* const speck_p = ptr + pos;
const auto num_bitplanes = speck_int_get_num_bitplanes(speck_p);
if (num_bitplanes <= 8)
m_uint_flag = UINTType::UINT8;
else if (num_bitplanes <= 16)
m_uint_flag = UINTType::UINT16;
else if (num_bitplanes <= 32)
m_uint_flag = UINTType::UINT32;
else
m_uint_flag = UINTType::UINT64;
m_instantiate_int_vec();
m_instantiate_decoder();
auto speck_suppose_len =
std::visit([speck_p](auto&& dec) { return dec->get_stream_full_len(speck_p); }, m_decoder);
auto speck_len = std::min(static_cast<size_t>(speck_suppose_len), remaining_len);
std::visit([speck_p, speck_len](auto&& dec) { return dec->use_bitstream(speck_p, speck_len); },
m_decoder);
pos += speck_len;
assert(pos <= len);
m_has_outlier = false;
if (pos < len) {
const uint8_t* const out_p = ptr + pos;
remaining_len = len - pos;
if (remaining_len >= SPECK_INT<uint8_t>::header_size) {
auto suppose_len = m_out_coder.get_stream_full_len(out_p);
assert(suppose_len >= remaining_len);
if (remaining_len == suppose_len) {
auto rtn = m_out_coder.use_bitstream(out_p, suppose_len);
if (rtn != RTNType::Good)
return rtn;
m_has_outlier = true;
}
}
}
return RTNType::Good;
}
void sperr::SPECK_FLT::append_encoded_bitstream(vec8_type& buf) const
{
std::copy(m_condi_bitstream.cbegin(), m_condi_bitstream.cend(), std::back_inserter(buf));
if (!m_conditioner.is_constant(m_condi_bitstream[0])) {
std::visit([&buf](auto&& enc) { enc->append_encoded_bitstream(buf); }, m_encoder);
if (m_has_outlier)
m_out_coder.append_encoded_bitstream(buf);
}
}
auto sperr::SPECK_FLT::view_decoded_data() const -> const vecd_type&
{
return m_vals_d;
}
auto sperr::SPECK_FLT::release_decoded_data() -> vecd_type&&
{
return std::move(m_vals_d);
}
auto sperr::SPECK_FLT::release_hierarchy() -> std::vector<vecd_type>&&
{
return std::move(m_hierarchy);
}
auto sperr::SPECK_FLT::view_hierarchy() const -> const std::vector<vecd_type>&
{
return m_hierarchy;
}
void sperr::SPECK_FLT::set_psnr(double psnr)
{
assert(psnr > 0.0);
m_quality = psnr;
m_mode = CompMode::PSNR;
m_q = 0.0; m_has_outlier = false;
}
void sperr::SPECK_FLT::set_tolerance(double tol)
{
assert(tol > 0.0);
m_quality = tol;
m_mode = CompMode::PWE;
m_q = 0.0; m_has_outlier = false;
}
void sperr::SPECK_FLT::set_bitrate(double bpp)
{
assert(bpp > 0.0);
m_quality = bpp;
m_mode = CompMode::Rate;
m_q = 0.0; m_has_outlier = false;
}
#ifdef EXPERIMENTING
void sperr::SPECK_FLT::set_direct_q(double q)
{
assert(q > 0.0);
m_quality = q;
m_mode = CompMode::DirectQ;
m_q = 0.0; m_has_outlier = false;
}
#endif
void sperr::SPECK_FLT::set_dims(dims_type dims)
{
m_dims = dims;
}
auto sperr::SPECK_FLT::integer_len() const -> size_t
{
switch (m_uint_flag) {
case UINTType::UINT8:
assert(m_vals_ui.index() == 0);
assert(m_encoder.index() == 0 || m_decoder.index() == 0);
return sizeof(uint8_t);
case UINTType::UINT16:
assert(m_vals_ui.index() == 1);
assert(m_encoder.index() == 1 || m_decoder.index() == 1);
return sizeof(uint16_t);
case UINTType::UINT32:
assert(m_vals_ui.index() == 2);
assert(m_encoder.index() == 2 || m_decoder.index() == 2);
return sizeof(uint32_t);
default:
assert(m_vals_ui.index() == 3);
assert(m_encoder.index() == 3 || m_decoder.index() == 3);
return sizeof(uint64_t);
}
}
void sperr::SPECK_FLT::m_instantiate_int_vec()
{
switch (m_uint_flag) {
case UINTType::UINT8:
if (m_vals_ui.index() != 0)
m_vals_ui = std::vector<uint8_t>();
break;
case UINTType::UINT16:
if (m_vals_ui.index() != 1)
m_vals_ui = std::vector<uint16_t>();
break;
case UINTType::UINT32:
if (m_vals_ui.index() != 2)
m_vals_ui = std::vector<uint32_t>();
break;
default:
if (m_vals_ui.index() != 3)
m_vals_ui = std::vector<uint64_t>();
}
}
auto sperr::SPECK_FLT::m_estimate_mse_midtread(double q) const -> double
{
assert(!m_vals_d.empty());
const auto len = m_vals_d.size();
const size_t stride_size = 4096;
const size_t num_strides = len / stride_size;
auto tmp_buf = vecd_type(num_strides + 1);
const auto rcp_q = 1.0 / q;
for (size_t i = 0; i < num_strides; i++) {
const auto beg = m_vals_d.cbegin() + i * stride_size;
tmp_buf[i] = std::accumulate(beg, beg + stride_size, 0.0, [q, rcp_q](auto init, auto v) {
auto diff = std::fma(-q, std::rint(v * rcp_q), v);
return init + diff * diff;
});
}
tmp_buf[num_strides] = 0.0;
tmp_buf[num_strides] = std::accumulate(m_vals_d.cbegin() + num_strides * stride_size,
m_vals_d.cend(), 0.0, [q, rcp_q](auto init, auto v) {
auto diff = std::fma(-q, std::rint(v * rcp_q), v);
return init + diff * diff;
});
const auto total_sum = std::accumulate(tmp_buf.cbegin(), tmp_buf.cend(), 0.0);
const auto mse = total_sum / static_cast<double>(len);
return mse;
}
auto sperr::SPECK_FLT::m_estimate_q(double param, bool high_prec) const -> double
{
switch (m_mode) {
case CompMode::PSNR: {
const auto t_mse = (param * param) * std::pow(10.0, -m_quality / 10.0);
auto q = 2.0 * std::sqrt(t_mse * 3.0);
while (m_estimate_mse_midtread(q) > t_mse)
q /= std::exp2(0.25); return q;
}
case CompMode::PWE:
return m_quality * 1.5;
case CompMode::Rate:
if (!high_prec) {
return param / static_cast<double>(std::numeric_limits<uint32_t>::max());
}
else {
return param / 0x1.fffffffffffffp52;
}
#ifdef EXPERIMENTING
case CompMode::DirectQ:
return m_quality;
#endif
default:
return 0.0;
}
}
auto sperr::SPECK_FLT::m_midtread_quantize() -> RTNType
{
std::fesetround(FE_TONEAREST);
assert(FE_TONEAREST == std::fegetround());
assert(FLT_ROUNDS == 1);
auto maxd = *std::max_element(m_vals_d.cbegin(), m_vals_d.cend(),
[](auto a, auto b) { return std::abs(a) < std::abs(b); });
assert(m_q > 0.0);
auto maxdq = std::abs(maxd) / m_q;
if (!std::isfinite(maxdq) || (maxdq < std::numeric_limits<long long>::lowest()) || (maxdq > std::numeric_limits<long long>::max())) {
return RTNType::FE_Invalid;
}
auto maxll = std::llrint(maxdq);
if (maxll <= std::numeric_limits<uint8_t>::max())
m_uint_flag = UINTType::UINT8;
else if (maxll <= std::numeric_limits<uint16_t>::max())
m_uint_flag = UINTType::UINT16;
else if (maxll <= std::numeric_limits<uint32_t>::max())
m_uint_flag = UINTType::UINT32;
else
m_uint_flag = UINTType::UINT64;
m_instantiate_int_vec();
const auto total_vals = m_vals_d.size();
std::visit([total_vals](auto&& vec) { vec.resize(total_vals); }, m_vals_ui);
m_sign_array.resize(total_vals);
std::visit(
[&vals_d = m_vals_d, &signs = m_sign_array, q = m_q](auto&& vec) {
auto inv = 1.0 / q;
auto bits_x64 = vals_d.size() - vals_d.size() % 64;
for (size_t i = 0; i < bits_x64; i += 64) {
auto bits64 = uint64_t{0};
for (size_t j = 0; j < 64; j++) {
auto ll = std::llrint(vals_d[i + j] * inv);
bits64 |= uint64_t{ll >= 0} << j;
vec[i + j] = std::abs(ll);
}
signs.wlong(i, bits64);
}
for (size_t i = bits_x64; i < vals_d.size(); i++) {
auto ll = std::llrint(vals_d[i] * inv);
signs.wbit(i, (ll >= 0));
vec[i] = std::abs(ll);
}
},
m_vals_ui);
return RTNType::Good;
}
void sperr::SPECK_FLT::m_midtread_inv_quantize()
{
assert(m_sign_array.size() == std::visit([](auto&& vec) { return vec.size(); }, m_vals_ui));
assert(m_q > 0.0);
const auto tmpd = std::array<double, 2>{-1.0, 1.0};
m_vals_d.resize(m_sign_array.size());
std::visit(
[&vals_d = m_vals_d, &signs = m_sign_array, q = m_q, tmpd](auto&& vec) {
auto bits_x64 = vals_d.size() - vals_d.size() % 64;
for (size_t i = 0; i < bits_x64; i += 64) {
const auto bits64 = signs.rlong(i);
for (size_t j = 0; j < 64; j++) {
auto bit = (bits64 >> j) & uint64_t{1};
vals_d[i + j] = q * static_cast<double>(vec[i + j]) * tmpd[bit];
}
}
for (size_t i = bits_x64; i < vals_d.size(); i++)
vals_d[i] = q * static_cast<double>(vec[i]) * tmpd[signs.rbit(i)];
},
m_vals_ui);
}
auto sperr::SPECK_FLT::compress() -> RTNType
{
const auto total_vals = size_t(m_dims[0]) * m_dims[1] * m_dims[2];
if (m_vals_d.empty() || m_vals_d.size() != total_vals)
return RTNType::Error;
if (m_mode == sperr::CompMode::Unknown)
return RTNType::CompModeUnknown;
m_has_outlier = false;
m_condi_bitstream = m_conditioner.condition(m_vals_d, m_dims);
if (m_conditioner.is_constant(m_condi_bitstream[0]))
return RTNType::Good;
auto param_q = 0.0; switch (m_mode) {
case CompMode::PWE:
m_vals_orig.resize(total_vals);
std::copy(m_vals_d.cbegin(), m_vals_d.cend(), m_vals_orig.begin());
break;
case CompMode::PSNR: {
auto [min, max] = std::minmax_element(m_vals_d.cbegin(), m_vals_d.cend());
param_q = *max - *min;
break;
}
default:; }
m_cdf.take_data(std::move(m_vals_d), m_dims);
m_wavelet_xform();
m_vals_d = m_cdf.release_data();
if (m_mode == CompMode::Rate) {
auto itr = std::max_element(m_vals_d.cbegin(), m_vals_d.cend(),
[](auto a, auto b) { return std::abs(a) < std::abs(b); });
param_q = std::abs(*itr);
}
bool high_prec = false;
FIXED_RATE_HIGH_PREC_LABEL:
m_q = m_estimate_q(param_q, high_prec);
assert(m_q > 0.0);
m_conditioner.save_q(m_condi_bitstream, m_q);
auto rtn = m_midtread_quantize();
if (rtn != RTNType::Good)
return rtn;
if (m_mode == CompMode::PWE) {
m_midtread_inv_quantize();
rtn = m_cdf.take_data(std::move(m_vals_d), m_dims);
if (rtn != RTNType::Good)
return rtn;
m_inverse_wavelet_xform(false); m_vals_d = m_cdf.release_data();
auto LOS = std::vector<Outlier>();
LOS.reserve(total_vals / 20); for (size_t i = 0; i < total_vals; i++) {
auto diff = m_vals_orig[i] - m_vals_d[i];
if (std::abs(diff) > m_quality)
LOS.emplace_back(i, diff);
}
if (LOS.empty())
m_has_outlier = false;
else {
m_has_outlier = true;
m_out_coder.set_length(total_vals);
m_out_coder.set_tolerance(m_quality);
m_out_coder.use_outlier_list(std::move(LOS));
rtn = m_out_coder.encode();
if (rtn != RTNType::Good)
return rtn;
}
}
m_instantiate_encoder();
if (m_mode == CompMode::Rate) {
auto budget = static_cast<size_t>(m_quality * double(total_vals)); std::visit([budget](auto&& encoder) { encoder->set_budget(budget); }, m_encoder);
}
std::visit([&dims = m_dims](auto&& encoder) { encoder->set_dims(dims); }, m_encoder);
switch (m_uint_flag) {
case UINTType::UINT8:
assert(m_vals_ui.index() == 0);
assert(m_encoder.index() == 0);
rtn = std::get<0>(m_encoder)->use_coeffs(std::move(std::get<0>(m_vals_ui)),
std::move(m_sign_array));
break;
case UINTType::UINT16:
assert(m_vals_ui.index() == 1);
assert(m_encoder.index() == 1);
rtn = std::get<1>(m_encoder)->use_coeffs(std::move(std::get<1>(m_vals_ui)),
std::move(m_sign_array));
break;
case UINTType::UINT32:
assert(m_vals_ui.index() == 2);
assert(m_encoder.index() == 2);
rtn = std::get<2>(m_encoder)->use_coeffs(std::move(std::get<2>(m_vals_ui)),
std::move(m_sign_array));
break;
default:
assert(m_vals_ui.index() == 3);
assert(m_encoder.index() == 3);
rtn = std::get<3>(m_encoder)->use_coeffs(std::move(std::get<3>(m_vals_ui)),
std::move(m_sign_array));
}
if (rtn != RTNType::Good)
return rtn;
std::visit([](auto&& encoder) { encoder->encode(); }, m_encoder);
if (m_mode == CompMode::Rate && high_prec == false) {
assert(m_encoder.index() == 2);
auto budget = static_cast<size_t>(m_quality * double(total_vals));
auto actual = std::get<2>(m_encoder)->encoded_bitstream_len() * size_t{8};
if (actual < budget) {
high_prec = true;
goto FIXED_RATE_HIGH_PREC_LABEL;
}
}
return RTNType::Good;
}
auto sperr::SPECK_FLT::decompress(bool multi_res) -> RTNType
{
m_vals_d.clear();
std::visit([](auto&& vec) { vec.clear(); }, m_vals_ui);
m_sign_array.resize(0);
if (m_conditioner.is_constant(m_condi_bitstream[0])) {
auto rtn = m_conditioner.inverse_condition(m_vals_d, m_dims, m_condi_bitstream);
return rtn;
}
assert(m_q > 0.0);
std::visit([dims = m_dims](auto&& decoder) { decoder->set_dims(dims); }, m_decoder);
std::visit([](auto&& decoder) { decoder->decode(); }, m_decoder);
std::visit([&vec = m_vals_ui](auto&& dec) { vec = dec->release_coeffs(); }, m_decoder);
m_sign_array = std::visit([](auto&& dec) { return dec->release_signs(); }, m_decoder);
m_midtread_inv_quantize();
auto rtn = m_cdf.take_data(std::move(m_vals_d), m_dims);
if (rtn != RTNType::Good)
return rtn;
m_inverse_wavelet_xform(multi_res);
m_vals_d = m_cdf.release_data();
if (m_has_outlier) {
m_out_coder.set_length(m_dims[0] * m_dims[1] * m_dims[2]);
m_out_coder.set_tolerance(m_q / 1.5); rtn = m_out_coder.decode();
if (rtn != RTNType::Good)
return rtn;
const auto& recovered = m_out_coder.view_outlier_list();
for (auto out : recovered)
m_vals_d[out.pos] += out.err;
}
rtn = m_conditioner.inverse_condition(m_vals_d, m_dims, m_condi_bitstream);
if (rtn != RTNType::Good)
return rtn;
if (multi_res) {
auto resolutions = sperr::coarsened_resolutions(m_dims);
if (m_hierarchy.size() != resolutions.size())
return RTNType::Error;
for (size_t h = 0; h < m_hierarchy.size(); h++) {
const auto& res = resolutions[h];
if (m_hierarchy[h].size() != res[0] * res[1] * res[2])
return RTNType::Error;
else
m_conditioner.inverse_condition(m_hierarchy[h], res, m_condi_bitstream);
}
}
return RTNType::Good;
}