Team Ai
Apppublic

tomaszki/PythonFileCompressor

sourceHugging Faceupdated 3y agoView on Hugging Face
1likes
numpyAc_backend.cpp337 linesDownload Raw Back to backend
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