#ifndef __ENCODE_HPP__
#define __ENCODE_HPP__
#define CODE_VALUE_BITS 32
#include <map>
#include <iterator>
#include "io.hpp"
using namespace std;
inline void put_bit(writer &w, uint64_t &encoding_bits, char bit) {
write_bits(w, bit, 1);
encoding_bits++;
}
inline void put_bit_plus_pending(writer &w, uint64_t &encoding_bits, bool bit, int& pending_bits)
{
put_bit(w, encoding_bits, bit);
for ( int i = 0 ; i < pending_bits ; i++ )
put_bit(w, encoding_bits, !bit);
pending_bits = 0;
}
uint64_t encode(writer &w, vector<uint64_t>& rle) {
constexpr uint64_t MAX_CODE = (((uint64_t)1) << CODE_VALUE_BITS)-1;
constexpr uint64_t ONE_FOURTH = (MAX_CODE + ((uint64_t)1))/4;
constexpr uint64_t ONE_HALF = ONE_FOURTH*2;
constexpr uint64_t THREE_FOURTHS = ONE_FOURTH*3;
std::map<uint64_t, pair<uint64_t, uint64_t> > frequencies;
for (uint64_t i = 0; i < rle.size(); ++i)
++frequencies[rle[i]].first;
uint64_t count = 0;
for (std::map<uint64_t, pair<uint64_t, uint64_t> >::iterator it = frequencies.begin(); it != frequencies.end(); ++it) {
(it->second).second = count;
count += (it->second).first;
}
uint64_t encoding_bits = 0;
uint64_t dict_size = frequencies.size();
write_bits(w, dict_size, sizeof(uint64_t)*8);
encoding_bits += sizeof(uint64_t)*8;
for (std::map<uint64_t, pair<uint64_t, uint64_t> >::iterator it = frequencies.begin(); it != frequencies.end(); ++it) {
uint64_t key = it->first;
uint64_t freq = (it->second).first;
uint8_t key_len = 0;
uint64_t key_copy = key;
while (key_copy) {
key_copy >>= 1;
key_len++;
}
key_len = max(1, key_len); write_bits(w, key_len, 6);
write_bits(w, key, key_len);
uint8_t freq_len = 0;
uint64_t freq_copy = freq;
while (freq_copy) {
freq_copy >>= 1;
freq_len++;
}
freq_len = max(1, freq_len); write_bits(w, freq_len, 6);
write_bits(w, freq, freq_len);
encoding_bits += 6 + key_len + 6 + freq_len;
}
uint64_t n_symbols = rle.size();
write_bits(w, n_symbols, sizeof(uint64_t)*8);
encoding_bits += sizeof(uint64_t)*8;
int pending_bits = 0;
uint64_t low = 0;
uint64_t high = MAX_CODE;
uint64_t rle_pos = 0;
for ( ; ; ) {
uint64_t c = rle[rle_pos];
rle_pos++;
uint64_t phigh = frequencies[c].second + frequencies[c].first;
uint64_t plow = frequencies[c].second;
uint64_t range = high - low + 1;
high = low + (range * phigh / n_symbols) - 1;
low = low + (range * plow / n_symbols);
for ( ; ; ) {
if ( high < ONE_HALF )
put_bit_plus_pending(w, encoding_bits, 0, pending_bits);
else if ( low >= ONE_HALF )
put_bit_plus_pending(w, encoding_bits, 1, pending_bits);
else if ( low >= ONE_FOURTH && high < THREE_FOURTHS ) {
pending_bits++;
low -= ONE_FOURTH;
high -= ONE_FOURTH;
} else
break;
high <<= 1;
high++;
low <<= 1;
high &= MAX_CODE;
low &= MAX_CODE;
}
if (rle_pos == n_symbols)
break;
}
pending_bits++;
if ( low < ONE_FOURTH )
put_bit_plus_pending(w, encoding_bits, 0, pending_bits);
else
put_bit_plus_pending(w, encoding_bits, 1, pending_bits);
write_bits(w, 0UL, CODE_VALUE_BITS-2); encoding_bits += CODE_VALUE_BITS-2;
return encoding_bits;
}
#endif