52#include <unordered_map>
69 std::is_same_v<T, uint8_t> ||
70 std::is_same_v<T, uint16_t> ||
71 std::is_same_v<T, uint32_t>,
72 "BitpackStage: T must be uint8_t, uint16_t, or uint32_t.");
79 void setInverse(
bool inv)
override { is_inverse_ = inv; }
80 bool isInverse()
const override {
return is_inverse_; }
96 if (nbits == 0 || nbits > 8 *
sizeof(T) || (nbits & (nbits - 1)) != 0)
97 throw std::invalid_argument(
98 "BitpackStage::setNBits: nbits must be a power of two "
99 "in [1, " + std::to_string(8 *
sizeof(T)) +
"], got "
100 + std::to_string(nbits));
103 uint8_t getNBits()
const {
return nbits_; }
113 T getBase()
const {
return base_; }
127 if (shift >= 8 *
sizeof(T))
128 throw std::invalid_argument(
129 "BitpackStage::setShift: shift must be in [0, "
130 + std::to_string(8 *
sizeof(T) - 1) +
"], got "
131 + std::to_string(shift));
134 uint8_t getShift()
const {
return shift_; }
150 bool isAutoDetect()
const {
return auto_detect_; }
158 bool isAutoBase()
const {
return auto_base_; }
168 bool isAutoShift()
const {
return auto_shift_; }
177 auto_base_ = auto_shift_ = auto_detect_ = enable;
184 const std::vector<void*>& inputs,
185 const std::vector<void*>& outputs,
186 const std::vector<size_t>& sizes
190 std::string
getName()
const override {
return "Bitpack"; }
191 size_t getNumInputs()
const override {
return 1; }
192 size_t getNumOutputs()
const override {
return 1; }
195 const std::vector<size_t>& input_sizes
197 if (input_sizes.empty())
return {0};
202 return {input_sizes[0]};
205 const size_t n = input_sizes[0] /
sizeof(T);
206 return {(n * nbits_ + 7) / 8};
210 const size_t max_elems = (input_sizes[0] * 8 + nbits_ - 1) / nbits_;
211 return {max_elems *
sizeof(T)};
215 std::unordered_map<std::string, size_t>
217 return {{
"output", actual_output_size_}};
221 return (index == 0) ? actual_output_size_ : 0;
227 return static_cast<uint16_t
>(StageType::BITPACK);
241 size_t , uint8_t* buf,
size_t max_size
243 if (max_size < 15)
return 0;
244 buf[0] =
static_cast<uint8_t
>(dataTypeOf<T>());
246 std::memcpy(buf + 2, &num_elements_,
sizeof(uint64_t));
248 const uint32_t base32 =
static_cast<uint32_t
>(base_);
249 std::memcpy(buf + 11, &base32,
sizeof(uint32_t));
256 if (size >= 2) nbits_ = buf[1];
257 if (size >= 10) std::memcpy(&num_elements_, buf + 2,
sizeof(uint64_t));
262 std::memcpy(&base32, buf + 11,
sizeof(uint32_t));
263 base_ =
static_cast<T
>(base32);
273 saved_nbits_ = nbits_;
274 saved_num_elements_ = num_elements_;
275 saved_output_size_ = actual_output_size_;
276 saved_shift_ = shift_;
280 void restoreState()
override {
281 nbits_ = saved_nbits_;
282 num_elements_ = saved_num_elements_;
283 actual_output_size_ = saved_output_size_;
284 shift_ = saved_shift_;
291 return !(auto_detect_ || auto_base_ || auto_shift_);
295 bool is_inverse_ =
false;
296 bool auto_detect_ =
false;
297 bool auto_base_ =
false;
298 bool auto_shift_ =
false;
299 uint8_t nbits_ = 8 *
sizeof(T);
302 uint64_t num_elements_ = 0;
303 size_t actual_output_size_ = 0;
306 uint8_t saved_nbits_ = 8 *
sizeof(T);
307 uint8_t saved_shift_ = 0;
308 T saved_base_ = T(0);
309 uint64_t saved_num_elements_ = 0;
310 size_t saved_output_size_ = 0;
313 static constexpr DataType dataTypeOf() {
314 if (std::is_same_v<U, uint8_t>)
return DataType::UINT8;
315 if (std::is_same_v<U, uint16_t>)
return DataType::UINT16;
316 if (std::is_same_v<U, uint32_t>)
return DataType::UINT32;
317 return DataType::UINT8;
321extern template class BitpackStage<uint8_t>;
322extern template class BitpackStage<uint16_t>;
323extern template class BitpackStage<uint32_t>;
Definition bitpack_stage.h:66
void setAdaptive(bool enable)
Definition bitpack_stage.h:176
void setAutoShift(bool enable)
Definition bitpack_stage.h:167
std::vector< size_t > estimateOutputSizes(const std::vector< size_t > &input_sizes) const override
Definition bitpack_stage.h:194
void saveState() override
Definition bitpack_stage.h:272
bool isGraphCompatible() const override
Definition bitpack_stage.h:290
void execute(fz::stream_t stream, MemoryPool *pool, const std::vector< void * > &inputs, const std::vector< void * > &outputs, const std::vector< size_t > &sizes) override
size_t getActualOutputSize(int index) const override
Definition bitpack_stage.h:220
uint8_t getInputDataType(size_t) const override
Definition bitpack_stage.h:234
void setShift(uint8_t shift)
Definition bitpack_stage.h:126
std::string getName() const override
Definition bitpack_stage.h:190
void setAutoDetect(bool enable)
Definition bitpack_stage.h:149
uint16_t getStageTypeId() const override
Definition bitpack_stage.h:226
std::unordered_map< std::string, size_t > getActualOutputSizesByName() const override
Definition bitpack_stage.h:216
size_t getMaxHeaderSize(size_t) const override
Definition bitpack_stage.h:267
void setInverse(bool inv) override
Definition bitpack_stage.h:79
uint8_t getOutputDataType(size_t) const override
Definition bitpack_stage.h:231
void deserializeHeader(const uint8_t *buf, size_t size) override
Definition bitpack_stage.h:253
void setAutoBase(bool enable)
Definition bitpack_stage.h:157
size_t serializeHeader(size_t, uint8_t *buf, size_t max_size) const override
Definition bitpack_stage.h:240
void setNBits(uint8_t nbits)
Definition bitpack_stage.h:95
void setBase(T base)
Definition bitpack_stage.h:112
DataType
Element data type identifiers used in buffer and stage descriptors.
Definition fzm_format.h:139
@ UNKNOWN
Byte-transparent stages: skip type checking at finalize()
Base class interface for all compression stages.
Backend-neutral GPU type aliases.