FZGPUModules 2.0
GPU-accelerated modular compression pipelines
Loading...
Searching...
No Matches
quant_adaptive_lorenzo_stage.h
Go to the documentation of this file.
1#pragma once
2
11#include "fused/lorenzo_quant/lorenzo_quant.h" // ErrorBoundMode, resolveApproxRelMode
12#include "stage/stage.h"
13#include "fzm_format.h"
14#include "backend/types.h"
16#include <cstdint>
17#include <cstring>
18#include <stdexcept>
19#include <string>
20#include <type_traits>
21#include <unordered_map>
22#include <vector>
23
24namespace fz {
25
36 uint8_t blocks_per_tile;
37 uint8_t enable_order2;
38 uint8_t enable_centering;
39 uint8_t eb_mode;
40 uint8_t reserved[3];
41 uint32_t num_elements;
42 double abs_error_bound_f64;
45
47 : coder_block_size(32), blocks_per_tile(8), enable_order2(1),
48 enable_centering(1), eb_mode(0), reserved{0, 0, 0}, num_elements(0),
49 abs_error_bound_f64(0.0), value_base_f64(0.0), user_error_bound(0.0f) {}
50};
51static_assert(sizeof(FusedQuantAdaptiveLorenzoConfig) <= FZM_STAGE_CONFIG_SIZE,
52 "FusedQuantAdaptiveLorenzoConfig must fit in FZM_STAGE_CONFIG_SIZE");
53
85template<typename T = int32_t>
87 static_assert(std::is_integral<T>::value && std::is_signed<T>::value,
88 "FusedQuantAdaptiveLorenzoStage requires a signed integer residual type");
89public:
90 struct Config {
91 uint32_t coder_block_size = 32;
92 uint32_t blocks_per_tile = 8;
93 bool enable_order2 = true;
94 bool enable_centering = true;
96 double error_bound = 1e-3;
103 float precomputed_value_base = 0.0f;
104 Config() = default;
105 };
106
107 FusedQuantAdaptiveLorenzoStage() { validate(); }
108 explicit FusedQuantAdaptiveLorenzoStage(const Config& config) : config_(config) { validate(); }
109
110 void setInverse(bool inv) override { is_inverse_ = inv; }
111 bool isInverse() const override { return is_inverse_; }
112
113 uint32_t getTileSize() const {
114 return config_.coder_block_size * config_.blocks_per_tile;
115 }
116
117 void setErrorBound(double eb) { config_.error_bound = eb; }
118 void setErrorBoundMode(ErrorBoundMode mode) {
119 config_.eb_mode = resolveApproxRelMode(mode, "FusedQuantAdaptiveLorenzoStage");
120 }
121 void setValueBase(float value_base) { config_.precomputed_value_base = value_base; }
122 double getErrorBound() const { return config_.error_bound; }
123 ErrorBoundMode getErrorBoundMode() const { return config_.eb_mode; }
126 double getComputedAbsErrorBound() const { return computed_abs_eb_; }
127
131 if (!decl.valid() || !decl.additive ||
132 (decl.kind != EncodingOracleKind::PlainFixedRateBitpack &&
133 decl.kind != EncodingOracleKind::AdaptiveFixedRateBitpack) ||
134 decl.input_data_type != static_cast<uint8_t>(getElementDataType()) ||
135 decl.unit_elems != config_.coder_block_size) {
136 return false;
137 }
138 bound_oracle_ = decl;
139 has_bound_oracle_ = true;
140 return true;
141 }
142
143 bool hasBoundEncodingOracle() const { return has_bound_oracle_; }
144 EncodingOracleKind getBoundEncodingOracleKind() const {
145 return has_bound_oracle_ ? bound_oracle_.kind
146 : EncodingOracleKind::PlainFixedRateBitpack;
147 }
148
149 FusionSpec getFusionSpec() const override {
150 if (is_inverse_ || !has_bound_oracle_) return {};
151 return FusionSpec{FusionAccess::TileSelector, getTileSize(),
152 config_.coder_block_size};
153 }
154
155 std::vector<FusedAuxOutputDecl> getFusedAuxOutputs() const override {
156 if (!getFusionSpec().fusable()) return {};
157 return {
159 static_cast<uint8_t>(DataType::UINT8), getTileSize(),
160 2u, 0u},
162 static_cast<uint8_t>(getElementDataType()), getTileSize(),
163 0u, 1u},
164 };
165 }
166
168 fz::stream_t stream,
169 MemoryPool* pool,
170 const std::vector<void*>& inputs,
171 const std::vector<void*>& outputs,
172 const std::vector<size_t>& sizes
173 ) override;
174
175 std::string getName() const override { return "FusedQuantAdaptiveLorenzo"; }
176 size_t getNumInputs() const override { return is_inverse_ ? 3 : 1; }
177 size_t getNumOutputs() const override { return is_inverse_ ? 1 : 3; }
178
179 std::vector<std::string> getOutputNames() const override {
180 return {"output", "modes", "means"};
181 }
182
183 ~FusedQuantAdaptiveLorenzoStage() override { releaseScratch(); }
184
185 void postStreamSync(fz::stream_t stream) override;
186
190 bool isGraphCompatible() const override { return false; }
191
193 const std::vector<size_t>& input_sizes
194 ) const override;
195
196 std::vector<size_t> estimateOutputSizes(
197 const std::vector<size_t>& input_sizes
198 ) const override {
199 if (input_sizes.empty()) return is_inverse_ ? std::vector<size_t>{0}
200 : std::vector<size_t>{0, 0, 0};
201 if (is_inverse_) return {input_sizes[0] / sizeof(T) * sizeof(float)};
202 const size_t n = input_sizes[0] / sizeof(float);
203 const size_t tiles = numTiles(n);
204 return {n * sizeof(T), (tiles + 3) / 4, tiles * sizeof(T)};
205 }
206
207 void saveState() override { saved_output_sizes_ = actual_output_sizes_; }
208 void restoreState() override {
209 if (!saved_output_sizes_.empty()) actual_output_sizes_ = saved_output_sizes_;
210 }
211
212 std::unordered_map<std::string, size_t>
213 getActualOutputSizesByName() const override {
214 auto names = getOutputNames();
215 std::unordered_map<std::string, size_t> r;
216 for (size_t i = 0; i < names.size() && i < actual_output_sizes_.size(); ++i)
217 r[names[i]] = actual_output_sizes_[i];
218 return r;
219 }
220
221 size_t getActualOutputSize(int index) const override {
222 return (index >= 0 && index < static_cast<int>(actual_output_sizes_.size()))
223 ? actual_output_sizes_[index] : 0;
224 }
225
226 void setFusedSideOutput(int output_index, size_t bytes) override {
227 if (actual_output_sizes_.size() < 3) actual_output_sizes_.resize(3, 0);
228 if (output_index == 1 || output_index == 2)
229 actual_output_sizes_[static_cast<size_t>(output_index)] = bytes;
230 }
231
232 uint16_t getStageTypeId() const override {
233 return static_cast<uint16_t>(StageType::FUSED_QUANT_ADAPTIVE_LORENZO);
234 }
235
236 uint8_t getOutputDataType(size_t output_index) const override {
237 if (is_inverse_) return static_cast<uint8_t>(DataType::FLOAT32);
238 return static_cast<uint8_t>(output_index == 1 ? DataType::UINT8
239 : getElementDataType());
240 }
241 uint8_t getInputDataType(size_t input_index) const override {
242 if (!is_inverse_) return static_cast<uint8_t>(DataType::FLOAT32);
243 return static_cast<uint8_t>(input_index == 1 ? DataType::UINT8
244 : getElementDataType());
245 }
246
247 size_t serializeHeader(size_t /*output_index*/, uint8_t* buf, size_t max_size) const override {
248 if (max_size < sizeof(FusedQuantAdaptiveLorenzoConfig))
249 throw std::runtime_error("FusedQuantAdaptiveLorenzoStage: header buffer too small");
251 cfg.coder_block_size = static_cast<uint8_t>(config_.coder_block_size);
252 cfg.blocks_per_tile = static_cast<uint8_t>(config_.blocks_per_tile);
253 cfg.enable_order2 = config_.enable_order2 ? 1u : 0u;
254 cfg.enable_centering = config_.enable_centering ? 1u : 0u;
255 cfg.eb_mode = static_cast<uint8_t>(config_.eb_mode);
256 cfg.num_elements = static_cast<uint32_t>(num_elements_);
257 cfg.abs_error_bound_f64 = computed_abs_eb_;
258 cfg.value_base_f64 = computed_value_base_;
259 cfg.user_error_bound = static_cast<float>(config_.error_bound);
260 std::memcpy(buf, &cfg, sizeof(cfg));
261 return sizeof(cfg);
262 }
263
264 void deserializeHeader(const uint8_t* buf, size_t size) override {
265 if (size < sizeof(FusedQuantAdaptiveLorenzoConfig))
266 throw std::runtime_error("FusedQuantAdaptiveLorenzoStage: header too small");
268 std::memcpy(&cfg, buf, sizeof(cfg));
269 config_.coder_block_size = cfg.coder_block_size;
270 config_.blocks_per_tile = cfg.blocks_per_tile;
271 config_.enable_order2 = (cfg.enable_order2 != 0);
272 config_.enable_centering = (cfg.enable_centering != 0);
273 config_.eb_mode = static_cast<ErrorBoundMode>(cfg.eb_mode);
274 num_elements_ = cfg.num_elements;
275 computed_abs_eb_ = cfg.abs_error_bound_f64;
276 computed_value_base_ = cfg.value_base_f64;
277 validate();
278 }
279
280 size_t getMaxHeaderSize(size_t /*output_index*/) const override {
281 return sizeof(FusedQuantAdaptiveLorenzoConfig);
282 }
283
284private:
285 Config config_;
286 EncodingOracleDecl bound_oracle_;
287 bool has_bound_oracle_ = false;
288 bool is_inverse_ = false;
289 size_t num_elements_ = 0;
292 double computed_abs_eb_ = 0.0;
293 double computed_value_base_ = 0.0;
294 std::vector<size_t> actual_output_sizes_{0, 0, 0};
295 std::vector<size_t> saved_output_sizes_;
296
297 // Forward scratch — identical role to AdaptiveLorenzoStage's, reused for
298 // the inverse's offset recomputation too (same pattern as that class).
299 uint8_t* d_modes_dense_ = nullptr;
300 T* d_means_dense_ = nullptr;
301 uint32_t* d_flags_ = nullptr;
302 uint32_t* d_offsets_ = nullptr;
303 size_t scratch_tiles_ = 0;
304 MemoryPool* scratch_pool_ = nullptr;
305 size_t pending_tiles_ = 0;
306
307 // Inverse-only scratch: the reconstructed integer codes before the final
308 // dequant pass (this stage's inverse always ends in float, unlike
309 // AdaptiveLorenzoStage's, which stops at T).
310 T* d_inverse_codes_ = nullptr;
311 size_t inverse_codes_elems_ = 0;
312 MemoryPool* inverse_scratch_pool_ = nullptr;
313
314 size_t ensureScratch(size_t num_tiles, MemoryPool* pool, fz::stream_t stream);
315 size_t ensureInverseScratch(size_t n, MemoryPool* pool, fz::stream_t stream);
316 void releaseScratch();
317
318 size_t numTiles(size_t n) const {
319 const size_t t = getTileSize();
320 return (n + t - 1) / t;
321 }
322
323 void validate() const {
324 if (config_.coder_block_size != 32)
325 throw std::invalid_argument(
326 "FusedQuantAdaptiveLorenzoStage: coder_block_size must be 32");
327 if (config_.blocks_per_tile < 1 || config_.blocks_per_tile > 32)
328 throw std::invalid_argument(
329 "FusedQuantAdaptiveLorenzoStage: blocks_per_tile must be in [1, 32], got "
330 + std::to_string(config_.blocks_per_tile));
331 }
332
333 static DataType getElementDataType() { return fused::dataTypeOf<T>(); }
334};
335
336extern template class FusedQuantAdaptiveLorenzoStage<int32_t>;
337
338} // namespace fz
Per-tile adaptive multi-order Lorenzo predictor with centering. Lossless.
Definition quant_adaptive_lorenzo_stage.h:86
void postStreamSync(fz::stream_t stream) override
size_t getActualOutputSize(int index) const override
Definition quant_adaptive_lorenzo_stage.h:221
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 setInverse(bool inv) override
Definition quant_adaptive_lorenzo_stage.h:110
FusionSpec getFusionSpec() const override
Definition quant_adaptive_lorenzo_stage.h:149
double getComputedAbsErrorBound() const
Definition quant_adaptive_lorenzo_stage.h:126
bool bindDownstreamEncodingOracle(const EncodingOracleDecl &decl) override
Definition quant_adaptive_lorenzo_stage.h:130
size_t estimateScratchBytes(const std::vector< size_t > &input_sizes) const override
size_t serializeHeader(size_t, uint8_t *buf, size_t max_size) const override
Definition quant_adaptive_lorenzo_stage.h:247
uint8_t getInputDataType(size_t input_index) const override
Definition quant_adaptive_lorenzo_stage.h:241
bool isGraphCompatible() const override
Definition quant_adaptive_lorenzo_stage.h:190
void deserializeHeader(const uint8_t *buf, size_t size) override
Definition quant_adaptive_lorenzo_stage.h:264
std::string getName() const override
Definition quant_adaptive_lorenzo_stage.h:175
std::unordered_map< std::string, size_t > getActualOutputSizesByName() const override
Definition quant_adaptive_lorenzo_stage.h:213
std::vector< std::string > getOutputNames() const override
Definition quant_adaptive_lorenzo_stage.h:179
std::vector< size_t > estimateOutputSizes(const std::vector< size_t > &input_sizes) const override
Definition quant_adaptive_lorenzo_stage.h:196
uint8_t getOutputDataType(size_t output_index) const override
Definition quant_adaptive_lorenzo_stage.h:236
std::vector< FusedAuxOutputDecl > getFusedAuxOutputs() const override
Definition quant_adaptive_lorenzo_stage.h:155
void saveState() override
Definition quant_adaptive_lorenzo_stage.h:207
uint16_t getStageTypeId() const override
Definition quant_adaptive_lorenzo_stage.h:232
void setFusedSideOutput(int output_index, size_t bytes) override
Definition quant_adaptive_lorenzo_stage.h:226
size_t getMaxHeaderSize(size_t) const override
Definition quant_adaptive_lorenzo_stage.h:280
Definition mempool.h:82
Definition stage.h:31
Compile-time C++ type -> DataType enum mapping, shared by the fused stages that dispatch on multiple ...
FZM binary file format definitions — structs, enums, and helpers.
Fused Lorenzo predictor and quantizer stage.
Definition dag.h:24
ErrorBoundMode
Definition lorenzo_quant.h:42
@ ABS
Absolute error bound.
ErrorBoundMode resolveApproxRelMode(ErrorBoundMode mode, const char *stage_name)
Definition lorenzo_quant.h:63
constexpr size_t FZM_STAGE_CONFIG_SIZE
Per-stage serialized config slot (bytes)
Definition fzm_format.h:65
@ FUSED_QUANT_ADAPTIVE_LORENZO
AdaptiveLorenzo with the upstream linear Quantizer fused into its forward kernel (FSZ "M1" partial fu...
DataType
Element data type identifiers used in buffer and stage descriptors.
Definition fzm_format.h:142
@ CompactedElements
runtime count * sizeof(element)
@ FixedBitsPerUnit
ceil(num_units * bits_per_unit / 8)
EncodingOracleKind
Registered exact encoded-size policies used by an upstream adaptive stage for an algorithmic mode dec...
Definition fusion.h:90
Base class interface for all compression stages.
Host-side declaration of a local, exact encoded-size oracle.
Definition fusion.h:104
uint8_t input_data_type
DataType value; 0xFF = unknown.
Definition fusion.h:109
Definition fusion.h:131
Definition quant_adaptive_lorenzo_stage.h:34
uint8_t coder_block_size
Fixed at 32.
Definition quant_adaptive_lorenzo_stage.h:35
uint8_t eb_mode
ErrorBoundMode cast to uint8_t.
Definition quant_adaptive_lorenzo_stage.h:39
double value_base_f64
NOA/PREL scan result; 0 for ABS.
Definition quant_adaptive_lorenzo_stage.h:43
float user_error_bound
Original config_.error_bound, narrow copy (debug only).
Definition quant_adaptive_lorenzo_stage.h:44
A stage's fusion contract. Stages that can participate in a fused kernel override Stage::getFusionSpe...
Definition fusion.h:51
Backend-neutral GPU type aliases.