FZGPUModules 2.0
GPU-accelerated modular compression pipelines
Loading...
Searching...
No Matches
szx_stage.h
Go to the documentation of this file.
1#pragma once
2
23#include "stage/stage.h"
24#include "fzm_format.h"
25#include "backend/types.h"
26#include <cstdint>
27#include <cstring>
28#include <stdexcept>
29#include <string>
30#include <type_traits>
31#include <unordered_map>
32#include <vector>
33
34namespace fz {
35
41enum class SZxErrorMode : uint8_t { ABS = 0, NOA = 2 };
42
46struct SZxConfig {
48 uint8_t eb_mode;
49 uint8_t _pad[2];
50 uint32_t block_size;
51 uint64_t num_elements;
52 double error_bound;
53 double value_base;
54
55 SZxConfig()
56 : data_type(DataType::FLOAT32), eb_mode(0), _pad{},
57 block_size(128), num_elements(0), error_bound(0.0), value_base(0.0) {}
58};
59static_assert(sizeof(SZxConfig) <= FZM_STAGE_CONFIG_SIZE,
60 "SZxConfig must fit in FZM_STAGE_CONFIG_SIZE");
61
84template<typename T>
85class SZxStage : public Stage {
86 static_assert(std::is_same_v<T, float> || std::is_same_v<T, double>,
87 "SZxStage: T must be float or double.");
88public:
89 SZxStage() = default;
90 ~SZxStage() override;
91
92 // ── Stage control ──────────────────────────────────────────────────────
93 void setInverse(bool inv) override { is_inverse_ = inv; }
94 bool isInverse() const override { return is_inverse_; }
95 // Forward defers its data-dependent compressed-size readback to
96 // postStreamSync() (see AdaptiveBitpackStage), so ABS forward is
97 // graph-capturable. NOA forward needs a range reduce + host read inside
98 // execute(), so it is not; the inverse keeps a per-execute layout either way.
99 bool isGraphCompatible() const override {
100 return !is_inverse_ && eb_mode_ != SZxErrorMode::NOA;
101 }
102
103 void setBlockSize(uint32_t n) {
104 if (n == 0 || n > 4096)
105 throw std::invalid_argument(
106 "SZxStage::setBlockSize: n must be in [1, 4096], got "
107 + std::to_string(n));
108 block_size_ = n;
109 }
110 uint32_t getBlockSize() const { return block_size_; }
111
112 void setErrorBound(double eb) { user_eb_ = eb; }
113 double getErrorBound() const { return user_eb_; }
114 void setErrorMode(SZxErrorMode m) { eb_mode_ = m; }
115 SZxErrorMode getErrorMode() const { return eb_mode_; }
116
119 double getConstantBlockFraction() const { return const_block_frac_; }
120
121 // ── Execution ──────────────────────────────────────────────────────────
122 void execute(fz::stream_t stream, MemoryPool* pool,
123 const std::vector<void*>& inputs,
124 const std::vector<void*>& outputs,
125 const std::vector<size_t>& sizes) override;
126
127 void postStreamSync(fz::stream_t stream) override;
128
129private:
133 double resolveAbsEb(fz::stream_t stream, MemoryPool* pool, const T* d_in, size_t n);
134
135public:
136
137 // ── Metadata ───────────────────────────────────────────────────────────
138 std::string getName() const override { return "SZx"; }
139 size_t getNumInputs() const override { return 1; }
140 size_t getNumOutputs() const override { return 1; }
141
142 std::vector<size_t> estimateOutputSizes(
143 const std::vector<size_t>& input_sizes) const override;
145 const std::vector<size_t>& input_sizes) const override;
146
147 std::unordered_map<std::string, size_t>
148 getActualOutputSizesByName() const override {
149 return {{"output", actual_output_size_}};
150 }
151 size_t getActualOutputSize(int index) const override {
152 return (index == 0) ? actual_output_size_ : 0;
153 }
154 std::vector<std::string> getRunNotes() const override;
155
156 uint16_t getStageTypeId() const override {
157 return static_cast<uint16_t>(StageType::SZX);
158 }
159 uint8_t getOutputDataType(size_t) const override {
160 return static_cast<uint8_t>(is_inverse_ ? getElementDataType()
161 : DataType::UINT8);
162 }
163 uint8_t getInputDataType(size_t) const override {
164 return static_cast<uint8_t>(is_inverse_ ? DataType::UINT8
165 : getElementDataType());
166 }
167
168 // ── Serialization ──────────────────────────────────────────────────────
169 size_t serializeHeader(size_t, uint8_t* buf, size_t max_size) const override {
170 if (max_size < sizeof(SZxConfig)) return 0;
171 SZxConfig cfg;
172 cfg.data_type = getElementDataType();
173 cfg.eb_mode = static_cast<uint8_t>(eb_mode_);
174 cfg.block_size = block_size_;
175 cfg.num_elements = static_cast<uint64_t>(num_elements_);
176 cfg.error_bound = abs_eb_;
177 cfg.value_base = value_base_;
178 std::memcpy(buf, &cfg, sizeof(cfg));
179 return sizeof(cfg);
180 }
181 void deserializeHeader(const uint8_t* buf, size_t size) override {
182 if (size < sizeof(SZxConfig))
183 throw std::runtime_error("SZxStage: header too small");
184 SZxConfig cfg;
185 std::memcpy(&cfg, buf, sizeof(cfg));
186 block_size_ = cfg.block_size ? cfg.block_size : 128u;
187 num_elements_ = static_cast<size_t>(cfg.num_elements);
188 eb_mode_ = static_cast<SZxErrorMode>(cfg.eb_mode);
189 abs_eb_ = cfg.error_bound;
190 value_base_ = cfg.value_base;
191 }
192 size_t getMaxHeaderSize(size_t) const override { return sizeof(SZxConfig); }
193
194 void saveState() override {
195 saved_ = {block_size_, num_elements_, actual_output_size_, abs_eb_, value_base_};
196 }
197 void restoreState() override {
198 block_size_ = saved_.block_size;
199 num_elements_ = saved_.num_elements;
200 actual_output_size_ = saved_.actual_size;
201 abs_eb_ = saved_.abs_eb;
202 value_base_ = saved_.value_base;
203 }
204
205 size_t getNumElements() const { return num_elements_; }
206
207private:
208 bool is_inverse_ = false;
209 uint32_t block_size_ = 128;
210 SZxErrorMode eb_mode_ = SZxErrorMode::ABS;
211 double user_eb_ = 1e-3;
212 double abs_eb_ = 0.0;
213 double value_base_ = 0.0;
214 size_t num_elements_ = 0;
215 size_t actual_output_size_= 0;
216 double const_block_frac_ = 0.0;
217
218 // Forward-path persistent scratch, kept alive across execute() so the
219 // compressed-size readback can be deferred to postStreamSync() (mirrors
220 // AdaptiveBitpackStage). Grown lazily; freed in the destructor.
221 uint32_t* d_block_cost_ = nullptr;
222 uint32_t* d_block_offset_ = nullptr;
223 size_t scratch_blocks_ = 0;
224 MemoryPool* scratch_pool_ = nullptr;
225 size_t fwd_num_blocks_ = 0;
226 size_t fwd_meta_bytes_ = 0;
227
228 struct Saved { uint32_t block_size; size_t num_elements; size_t actual_size;
229 double abs_eb; double value_base; };
230 Saved saved_{128, 0, 0, 0.0, 0.0};
231
232 static DataType getElementDataType() {
233 return std::is_same<T, float>::value ? DataType::FLOAT32 : DataType::FLOAT64;
234 }
235};
236
237extern template class SZxStage<float>;
238extern template class SZxStage<double>;
239
240} // namespace fz
Definition mempool.h:82
Definition szx_stage.h:85
size_t getActualOutputSize(int index) const override
Definition szx_stage.h:151
std::unordered_map< std::string, size_t > getActualOutputSizesByName() const override
Definition szx_stage.h:148
std::vector< size_t > estimateOutputSizes(const std::vector< size_t > &input_sizes) const override
bool isGraphCompatible() const override
Definition szx_stage.h:99
std::string getName() const override
Definition szx_stage.h:138
uint8_t getInputDataType(size_t) const override
Definition szx_stage.h:163
size_t estimateScratchBytes(const std::vector< size_t > &input_sizes) const override
void setInverse(bool inv) override
Definition szx_stage.h:93
std::vector< std::string > getRunNotes() const override
double getConstantBlockFraction() const
Definition szx_stage.h:119
size_t getMaxHeaderSize(size_t) const override
Definition szx_stage.h:192
void postStreamSync(fz::stream_t stream) override
uint8_t getOutputDataType(size_t) const override
Definition szx_stage.h:159
uint16_t getStageTypeId() const override
Definition szx_stage.h:156
void deserializeHeader(const uint8_t *buf, size_t size) override
Definition szx_stage.h:181
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
void saveState() override
Definition szx_stage.h:194
size_t serializeHeader(size_t, uint8_t *buf, size_t max_size) const override
Definition szx_stage.h:169
Definition stage.h:31
FZM binary file format definitions — structs, enums, and helpers.
Definition dag.h:24
constexpr size_t FZM_STAGE_CONFIG_SIZE
Per-stage serialized config slot (bytes)
Definition fzm_format.h:65
@ SZX
SZx ultrafast EB compressor: per-block constant/non-constant classification + fixed-length residuals ...
DataType
Element data type identifiers used in buffer and stage descriptors.
Definition fzm_format.h:142
SZxErrorMode
Definition szx_stage.h:41
Base class interface for all compression stages.
Definition szx_stage.h:46
uint8_t eb_mode
SZxErrorMode cast to uint8_t.
Definition szx_stage.h:48
uint32_t block_size
Elements per block (SZx default 128).
Definition szx_stage.h:50
double value_base
value_range used for NOA→ABS conversion; else 0.
Definition szx_stage.h:53
DataType data_type
FLOAT32 / FLOAT64 (1B).
Definition szx_stage.h:47
uint8_t _pad[2]
Must be zero.
Definition szx_stage.h:49
double error_bound
Absolute bound after mode conversion (f64-safe).
Definition szx_stage.h:52
uint64_t num_elements
Original element count (sizes the inverse).
Definition szx_stage.h:51
Backend-neutral GPU type aliases.