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
82template<typename T>
83class SZxStage : public Stage {
84 static_assert(std::is_same_v<T, float> || std::is_same_v<T, double>,
85 "SZxStage: T must be float or double.");
86public:
87 SZxStage() = default;
88 ~SZxStage() override;
89
90 // ── Stage control ──────────────────────────────────────────────────────
91 void setInverse(bool inv) override { is_inverse_ = inv; }
92 bool isInverse() const override { return is_inverse_; }
93 // Forward defers its data-dependent compressed-size readback to
94 // postStreamSync() (see AdaptiveBitpackStage), so ABS forward is
95 // graph-capturable. NOA forward needs a range reduce + host read inside
96 // execute(), so it is not; the inverse keeps a per-execute layout either way.
97 bool isGraphCompatible() const override {
98 return !is_inverse_ && eb_mode_ != SZxErrorMode::NOA;
99 }
100
101 void setBlockSize(uint32_t n) {
102 if (n == 0 || n > 4096)
103 throw std::invalid_argument(
104 "SZxStage::setBlockSize: n must be in [1, 4096], got "
105 + std::to_string(n));
106 block_size_ = n;
107 }
108 uint32_t getBlockSize() const { return block_size_; }
109
110 void setErrorBound(double eb) { user_eb_ = eb; }
111 double getErrorBound() const { return user_eb_; }
112 void setErrorMode(SZxErrorMode m) { eb_mode_ = m; }
113 SZxErrorMode getErrorMode() const { return eb_mode_; }
114
117 double getConstantBlockFraction() const { return const_block_frac_; }
118
119 // ── Execution ──────────────────────────────────────────────────────────
120 void execute(fz::stream_t stream, MemoryPool* pool,
121 const std::vector<void*>& inputs,
122 const std::vector<void*>& outputs,
123 const std::vector<size_t>& sizes) override;
124
125 void postStreamSync(fz::stream_t stream) override;
126
127private:
131 double resolveAbsEb(fz::stream_t stream, MemoryPool* pool, const T* d_in, size_t n);
132
133public:
134
135 // ── Metadata ───────────────────────────────────────────────────────────
136 std::string getName() const override { return "SZx"; }
137 size_t getNumInputs() const override { return 1; }
138 size_t getNumOutputs() const override { return 1; }
139
140 std::vector<size_t> estimateOutputSizes(
141 const std::vector<size_t>& input_sizes) const override;
143 const std::vector<size_t>& input_sizes) const override;
144
145 std::unordered_map<std::string, size_t>
146 getActualOutputSizesByName() const override {
147 return {{"output", actual_output_size_}};
148 }
149 size_t getActualOutputSize(int index) const override {
150 return (index == 0) ? actual_output_size_ : 0;
151 }
152 std::vector<std::string> getRunNotes() const override;
153
154 uint16_t getStageTypeId() const override {
155 return static_cast<uint16_t>(StageType::SZX);
156 }
157 uint8_t getOutputDataType(size_t) const override {
158 return static_cast<uint8_t>(is_inverse_ ? getElementDataType()
159 : DataType::UINT8);
160 }
161 uint8_t getInputDataType(size_t) const override {
162 return static_cast<uint8_t>(is_inverse_ ? DataType::UINT8
163 : getElementDataType());
164 }
165
166 // ── Serialization ──────────────────────────────────────────────────────
167 size_t serializeHeader(size_t, uint8_t* buf, size_t max_size) const override {
168 if (max_size < sizeof(SZxConfig)) return 0;
169 SZxConfig cfg;
170 cfg.data_type = getElementDataType();
171 cfg.eb_mode = static_cast<uint8_t>(eb_mode_);
172 cfg.block_size = block_size_;
173 cfg.num_elements = static_cast<uint64_t>(num_elements_);
174 cfg.error_bound = abs_eb_;
175 cfg.value_base = value_base_;
176 std::memcpy(buf, &cfg, sizeof(cfg));
177 return sizeof(cfg);
178 }
179 void deserializeHeader(const uint8_t* buf, size_t size) override {
180 if (size < sizeof(SZxConfig))
181 throw std::runtime_error("SZxStage: header too small");
182 SZxConfig cfg;
183 std::memcpy(&cfg, buf, sizeof(cfg));
184 block_size_ = cfg.block_size ? cfg.block_size : 128u;
185 num_elements_ = static_cast<size_t>(cfg.num_elements);
186 eb_mode_ = static_cast<SZxErrorMode>(cfg.eb_mode);
187 abs_eb_ = cfg.error_bound;
188 value_base_ = cfg.value_base;
189 }
190 size_t getMaxHeaderSize(size_t) const override { return sizeof(SZxConfig); }
191
192 void saveState() override {
193 saved_ = {block_size_, num_elements_, actual_output_size_, abs_eb_, value_base_};
194 }
195 void restoreState() override {
196 block_size_ = saved_.block_size;
197 num_elements_ = saved_.num_elements;
198 actual_output_size_ = saved_.actual_size;
199 abs_eb_ = saved_.abs_eb;
200 value_base_ = saved_.value_base;
201 }
202
203 size_t getNumElements() const { return num_elements_; }
204
205private:
206 bool is_inverse_ = false;
207 uint32_t block_size_ = 128;
208 SZxErrorMode eb_mode_ = SZxErrorMode::ABS;
209 double user_eb_ = 1e-3;
210 double abs_eb_ = 0.0;
211 double value_base_ = 0.0;
212 size_t num_elements_ = 0;
213 size_t actual_output_size_= 0;
214 double const_block_frac_ = 0.0;
215
216 // Forward-path persistent scratch, kept alive across execute() so the
217 // compressed-size readback can be deferred to postStreamSync() (mirrors
218 // AdaptiveBitpackStage). Grown lazily; freed in the destructor.
219 uint32_t* d_block_cost_ = nullptr;
220 uint32_t* d_block_offset_ = nullptr;
221 size_t scratch_blocks_ = 0;
222 MemoryPool* scratch_pool_ = nullptr;
223 size_t fwd_num_blocks_ = 0;
224 size_t fwd_meta_bytes_ = 0;
225
226 struct Saved { uint32_t block_size; size_t num_elements; size_t actual_size;
227 double abs_eb; double value_base; };
228 Saved saved_{128, 0, 0, 0.0, 0.0};
229
230 static DataType getElementDataType() {
231 return std::is_same<T, float>::value ? DataType::FLOAT32 : DataType::FLOAT64;
232 }
233};
234
235extern template class SZxStage<float>;
236extern template class SZxStage<double>;
237
238} // namespace fz
Definition mempool.h:82
Definition szx_stage.h:83
size_t getActualOutputSize(int index) const override
Definition szx_stage.h:149
std::unordered_map< std::string, size_t > getActualOutputSizesByName() const override
Definition szx_stage.h:146
std::vector< size_t > estimateOutputSizes(const std::vector< size_t > &input_sizes) const override
bool isGraphCompatible() const override
Definition szx_stage.h:97
std::string getName() const override
Definition szx_stage.h:136
uint8_t getInputDataType(size_t) const override
Definition szx_stage.h:161
size_t estimateScratchBytes(const std::vector< size_t > &input_sizes) const override
void setInverse(bool inv) override
Definition szx_stage.h:91
std::vector< std::string > getRunNotes() const override
double getConstantBlockFraction() const
Definition szx_stage.h:117
size_t getMaxHeaderSize(size_t) const override
Definition szx_stage.h:190
void postStreamSync(fz::stream_t stream) override
uint8_t getOutputDataType(size_t) const override
Definition szx_stage.h:157
uint16_t getStageTypeId() const override
Definition szx_stage.h:154
void deserializeHeader(const uint8_t *buf, size_t size) override
Definition szx_stage.h:179
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:192
size_t serializeHeader(size_t, uint8_t *buf, size_t max_size) const override
Definition szx_stage.h:167
Definition stage.h:31
FZM binary file format definitions — structs, enums, and helpers.
Definition dag.h:24
@ NOA
Value-range relative bound (norm-of-absolute).
@ ABS
Absolute error bound.
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:139
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.