tomaszki/PythonFileCompressor
1
1/**2 * COPYRIGHT 2020 ETH Zurich3 * BASED on4 *5 * https://marknelson.us/posts/2014/10/19/data-compression-with-arithmetic-coding.html6 */7 8#include <torch/extension.h>9 10#include <iostream>11#include <vector>12#include <tuple>13#include <fstream>14#include <algorithm>15#include <string>16#include <chrono>17#include <numeric>18#include <iterator>19 20#include <bitset>21 22using cdf_t = uint16_t;23 24/** Encapsulates a pointer to a CDF tensor */25struct cdf_ptr {26 cdf_t* data; // expected to be a N_sym x Lp matrix, stored in row major.27 const int N_sym; // Number of symbols stored by `data`.28 const int Lp; // == L+1, where L is the number of possible values a symbol can take.29 cdf_ptr(cdf_t* data,30 const int N_sym,31 const int Lp) : data(data), N_sym(N_sym), Lp(Lp) {};32};33 34 35 36/** Class to save output bit by bit to a byte string */37class OutCacheString {38private:39public:40 std::string out="";41 uint8_t cache=0;42 uint8_t count=0;43 void append(const int bit) {44 cache <<= 1;45 cache |= bit;46 count += 1;47 if (count == 8) {48 out.append(reinterpret_cast<const char *>(&cache), 1);49 count = 0;50 }51 }52 void flush() {53 if (count > 0) {54 for (int i = count; i < 8; ++i) {55 append(0);56 }57 assert(count==0);58 }59 }60 void append_bit_and_pending(const int bit, uint64_t &pending_bits) {61 append(bit);62 while (pending_bits > 0) {63 append(!bit);64 pending_bits -= 1;65 }66 }67};68 69/** Class to read byte string bit by bit */70class InCacheString {71private:72 const std::string in_;73 74public:75 explicit InCacheString(const std::string& in) : in_(in) {};76 77 uint8_t cache=0;78 uint8_t cached_bits=0;79 size_t in_ptr=0;80 81 void get(uint32_t& value) {82 83 if (cached_bits == 0) {84 if (in_ptr == in_.size()){ 85 value <<= 1;86 return;87 }88 /// Read 1 byte89 90 cache = (uint8_t) in_[in_ptr];91 in_ptr++;92 cached_bits = 8;93 }94 value <<= 1;95 value |= (cache >> (cached_bits - 1)) & 1;96 cached_bits--;97 }98 99 void initialize(uint32_t& value) {100 for (int i = 0; i < 32; ++i) {101 get(value);102 }103 }104};105 106 107//------------------------------------------------------------------------------108 109 110cdf_t binsearch(py::list &cdf, cdf_t target, cdf_t max_sym,111 const int offset) /* i * Lp */112{113 cdf_t left = 0;114 cdf_t right = max_sym + 1; // len(cdf) == max_sym + 2115 116 while (left + 1 < right) { // ?117 // Left and right will be < 0x10000 in practice,118 // so left+right fits in uint16_t.119 const auto m = static_cast<const cdf_t>((left + right) / 2);120 const auto v = cdf[offset + m].cast<cdf_t>();121 if (v < target) {122 left = m;123 } else if (v > target) {124 right = m;125 } else {126 return m;127 }128 }129 130 return left;131}132 133 134class decode135{136private:137 138public:139 int dataID=0;140 const int Lp;// To calculate offset141 const int max_symbol;142 uint32_t low = 0;143 uint32_t high = 0xFFFFFFFFU;144 const uint32_t c_count = 0x10000U;145 const int precision = 16;146 cdf_t sym_i = 0;147 uint32_t value = 0;148 InCacheString in_cache;149 decode(const std::string &in, const int&sysNumDim_):in_cache(in),Lp(sysNumDim_),max_symbol(sysNumDim_-2){150 in_cache.initialize(value);151 152 };153 154 int16_t decodeAsym(py::list cdf) {155 156 const uint64_t span = static_cast<uint64_t>(high) - static_cast<uint64_t>(low) + 1;157 // always < 0x10000 ???158 const uint16_t count = ((static_cast<uint64_t>(value) - static_cast<uint64_t>(low) + 1) * c_count - 1) / span;159 160 int offset = 0;161 162 sym_i = binsearch(cdf, count, (cdf_t)max_symbol, offset);163 164 const uint32_t c_low = cdf[offset + sym_i].cast<cdf_t>();165 const uint32_t c_high = sym_i == max_symbol ? 0x10000U : cdf[offset + sym_i + 1].cast<cdf_t>();166 167 high = (low - 1) + ((span * static_cast<uint64_t>(c_high)) >> precision);168 low = (low) + ((span * static_cast<uint64_t>(c_low)) >> precision);169 170 while (true) {171 if (low >= 0x80000000U || high < 0x80000000U) {172 low <<= 1;173 high <<= 1;174 high |= 1;175 176 in_cache.get(value);177 178 } else if (low >= 0x40000000U && high < 0xC0000000U) {179 /**180 * 0100 0000 ... <= value < 1100 0000 ...181 * <=>182 * 0100 0000 ... <= value <= 1011 1111 ...183 * <=>184 * value starts with 01 or 10.185 * 01 - 01 == 00 | 10 - 01 == 01186 * i.e., with shifts187 * 01A -> 0A or 10A -> 1A, i.e., discard 2SB as it's all the same while we are in188 * near convergence189 */190 low <<= 1;191 low &= 0x7FFFFFFFU; // make MSB 0192 high <<= 1;193 high |= 0x80000001U; // add 1 at the end, retain MSB = 1194 value -= 0x40000000U;195 196 in_cache.get(value);197 198 } else {199 break; 200 }201 }202 203 return (int16_t)sym_i;204 }205 206};207 208const void check_sym(const torch::Tensor& sym) {209 TORCH_CHECK(sym.sizes().size() == 1,210 "Invalid size for sym. Expected just 1 dim.")211}212 213/** Get an instance of the `cdf_ptr` struct. */214const struct cdf_ptr get_cdf_ptr(const torch::Tensor& cdf)215{216 TORCH_CHECK(!cdf.is_cuda(), "cdf must be on CPU!")217 const auto s = cdf.sizes();218 TORCH_CHECK(s.size() == 2, "Invalid size for cdf! Expected (N, Lp)")219 220 const int N_sym = s[0];221 const int Lp = s[1];222 const auto cdf_acc = cdf.accessor<int16_t, 2>();223 cdf_t* cdf_ptr = (uint16_t*)cdf_acc.data();224 225 const struct cdf_ptr res(cdf_ptr, N_sym, Lp);226 return res;227}228 229 230// -----------------------------------------------------------------------------231 232 233/** Encode symbols `sym` with CDF represented by `cdf_ptr`. NOTE: this is not exposted to python. */234py::bytes encode(235 const cdf_ptr& cdf_ptr,236 const torch::Tensor& sym){237 238 OutCacheString out_cache;239 240 uint32_t low = 0;241 uint32_t high = 0xFFFFFFFFU;242 uint64_t pending_bits = 0;243 244 const int precision = 16;245 246 const cdf_t* cdf = cdf_ptr.data;247 const int N_sym = cdf_ptr.N_sym;248 const int Lp = cdf_ptr.Lp;249 const int max_symbol = Lp - 2;250 251 auto sym_ = sym.accessor<int16_t, 1>();252 253 for (int i = 0; i < N_sym; ++i) {254 const int16_t sym_i = sym_[i];255 256 const uint64_t span = static_cast<uint64_t>(high) - static_cast<uint64_t>(low) + 1;257 258 const int offset = i * Lp;259 // Left boundary is at offset + sym_i260 const uint32_t c_low = cdf[offset + sym_i];261 // Right boundary is at offset + sym_i + 1, except for the `max_symbol`262 // For which we hardcode the maxvalue. So if e.g.263 // L == 4, it means that Lp == 5, and the allowed symbols are264 // {0, 1, 2, 3}. The max symbol is thus Lp - 2 == 3. It's probability265 // is then given by c_max - cdf[-2].266 const uint32_t c_high = sym_i == max_symbol ? 0x10000U : cdf[offset + sym_i + 1];267 268 high = (low - 1) + ((span * static_cast<uint64_t>(c_high)) >> precision);269 low = (low) + ((span * static_cast<uint64_t>(c_low)) >> precision);270 271 while (true) {272 if (high < 0x80000000U) {273 out_cache.append_bit_and_pending(0, pending_bits);274 low <<= 1;275 high <<= 1;276 high |= 1;277 } else if (low >= 0x80000000U) {278 out_cache.append_bit_and_pending(1, pending_bits);279 low <<= 1;280 high <<= 1;281 high |= 1;282 } else if (low >= 0x40000000U && high < 0xC0000000U) {283 pending_bits++;284 low <<= 1;285 low &= 0x7FFFFFFF;286 high <<= 1;287 high |= 0x80000001;288 } else {289 break;290 }291 }292 }293 294 pending_bits += 1;295 296 if (pending_bits) {297 if (low < 0x40000000U) {298 out_cache.append_bit_and_pending(0, pending_bits);299 } else {300 out_cache.append_bit_and_pending(1, pending_bits);301 }302 }303 304 out_cache.flush();305 306#ifdef VERBOSE307 std::chrono::steady_clock::time_point end= std::chrono::steady_clock::now();308 std::cout << "Time difference (sec) = " << (std::chrono::duration_cast<std::chrono::microseconds>(end - begin).count()) /1000000.0 <<std::endl;309#endif310 311 return py::bytes(out_cache.out);312}313 314 315/** See torchac.py */316py::bytes encode_cdf(317 const torch::Tensor& cdf, /* NHWLp, must be on CPU! */318 const torch::Tensor& sym)319{320 check_sym(sym);321 const auto cdf_ptr = get_cdf_ptr(cdf);322 return encode(cdf_ptr, sym);323}324 325 326 327 328PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {329 m.def("encode_cdf", &encode_cdf, "Encode from CDF");330 331 py::class_<decode>(m, "decode")332 .def(py::init([] (const std::string in, const int&sysNumDim_) {333 return new decode(in,sysNumDim_);334 }))335 .def("decodeAsym", &decode::decodeAsym);336}337 