32#include <unordered_map>
37class ANSStage :
public Stage {
46 static constexpr bool isSupportedOnBackend() {
47#if defined(FZGMOD_BACKEND_HIP) || defined(FZGMOD_BACKEND_SYCL)
55 ~ANSStage()
override =
default;
64 void setProbBits(uint8_t pb) { prob_bits_ = pb; }
65 uint8_t getProbBits()
const {
return prob_bits_; }
68 void setInverse(
bool inv)
override { is_inverse_ = inv; }
69 bool isInverse()
const override {
return is_inverse_; }
72 bool isGraphCompatible()
const override {
return false; }
75 size_t getRequiredInputAlignment()
const override {
return 4; }
84 void onFinalize(
size_t estimated_inlen, MemoryPool* pool)
override;
86 size_t estimateDeviceFootprintBytes(
size_t inlen)
const override;
88 size_t estimateScratchBytes(
const std::vector<size_t>& input_sizes)
const override;
94 const std::vector<void*>& inputs,
95 const std::vector<void*>& outputs,
96 const std::vector<size_t>& sizes
100 std::string getName()
const override {
return "ANS"; }
101 size_t getNumInputs()
const override {
return 1; }
102 size_t getNumOutputs()
const override {
return 1; }
104 std::vector<size_t> estimateOutputSizes(
105 const std::vector<size_t>& input_sizes
107 if (input_sizes.empty())
return {0};
111 return {input_sizes[0] * 2 + 8192};
115 return {original_bytes_ > 0 ? original_bytes_ : input_sizes[0]};
118 std::unordered_map<std::string, size_t>
119 getActualOutputSizesByName()
const override {
120 return {{
"output", actual_output_size_}};
123 size_t getActualOutputSize(
int index)
const override {
124 return (index == 0) ? actual_output_size_ : 0;
128 uint16_t getStageTypeId()
const override {
133 uint8_t getOutputDataType(
size_t )
const override {
136 uint8_t getInputDataType(
size_t )
const override {
141 size_t serializeHeader(
142 size_t , uint8_t* buf,
size_t max_size
144 if (max_size < 12)
return 0;
146 buf[1] = buf[2] = buf[3] = 0;
147 std::memcpy(buf + 4, &original_bytes_,
sizeof(uint64_t));
151 void deserializeHeader(
const uint8_t* buf,
size_t size)
override {
155 std::memcpy(&original_bytes_, buf + 4,
sizeof(uint64_t));
158 size_t getMaxHeaderSize(
size_t )
const override {
return 12; }
160 void saveState()
override {
161 saved_prob_bits_ = prob_bits_;
162 saved_original_bytes_ = original_bytes_;
163 saved_output_size_ = actual_output_size_;
166 void restoreState()
override {
167 prob_bits_ = saved_prob_bits_;
168 original_bytes_ = saved_original_bytes_;
169 actual_output_size_ = saved_output_size_;
173 bool is_inverse_ =
false;
174 uint8_t prob_bits_ = 10;
175 uint64_t original_bytes_ = 0;
176 size_t actual_output_size_ = 0;
179 size_t cap_bytes_ = 0;
183 uint32_t* d_temp_histogram_ =
nullptr;
184 void* d_table_ =
nullptr;
185 uint8_t* d_compressed_blocks_ =
nullptr;
186 uint32_t* d_compressed_words_ =
nullptr;
187 uint32_t* d_comp_words_prefix_ =
nullptr;
188 void* d_temp_prefix_sum_ =
nullptr;
189 uint32_t* d_decode_table_ =
nullptr;
194 uint8_t last_header_bytes_[32] = {};
197 int hist_grid_dim_ = 0;
198 int hist_block_dim_ = 0;
199 int hist_shmem_use_ = 0;
200 int hist_r_per_block_ = 0;
203 uint8_t saved_prob_bits_ = 10;
204 uint64_t saved_original_bytes_ = 0;
205 size_t saved_output_size_ = 0;
209 void initScratch(
size_t inlen, MemoryPool* pool);
216 void executeForward(fz::stream_t stream, MemoryPool* pool,
217 uint8_t* in, uint8_t* out,
size_t byte_size);
218 void executeInverse(fz::stream_t stream, MemoryPool* pool,
219 uint8_t* in, uint8_t* out);
223 static constexpr size_t kUncoalescedStride = 128 + 5120;
Definition algorithms.h:48
@ ANS
rANS entropy coder (GPU, via dietGPU)
@ UNKNOWN
Byte-transparent stages: skip type checking at finalize()
Base class interface for all compression stages.
Backend-neutral GPU type aliases.