FZGPUModules 2.0
GPU-accelerated modular compression pipelines
Loading...
Searching...
No Matches
ginterp_stage.h
Go to the documentation of this file.
1#pragma once
2
13#include "stage/stage.h"
14#include "fzm_format.h"
15#include "fused/lorenzo_quant/lorenzo_quant.h" // for ErrorBoundMode
16
17#include "backend/types.h"
18#include <array>
19#include <cmath>
20#include <cstdint>
21#include <memory>
22#include <cstring>
23#include <string>
24#include <stdexcept>
25#include <type_traits>
26#include <unordered_map>
27#include <vector>
28
29// Forward-declared in the public header so callers don't pull in the
30// cuSZ-Hi type subset. The full definition lives in cusz_type_subset.h,
31// included only inside the stage TU.
33
34namespace fz {
35
41 // ── identity / dims ─────────────────────────────────────────────────────
42 double error_bound;
44 uint32_t quant_radius;
45 uint32_t num_elements;
46 uint32_t outlier_count;
49 uint8_t ndim;
50 uint8_t eb_mode;
51 uint32_t dim_x;
52 uint32_t dim_y;
53 uint32_t dim_z;
54 uint32_t anchor_dim_x;
55 uint32_t anchor_dim_y;
56 uint32_t anchor_dim_z;
57 float user_eb;
58 float value_base;
59
60 // ── resolved INTERPOLATION_PARAMS (phase 2 auto-tune output, 39 B) ──────
61 // Both encoder and decoder must use the exact same intp_param values, so
62 // the resolved params are written here on compress and consumed on decompress.
63 // Layout chosen to mirror `INTERPOLATION_PARAMS` field-for-field; we don't
64 // memcpy the struct directly because its padding is implementation-defined.
65 double intp_alpha;
66 double intp_beta;
67 uint8_t intp_use_md[6];
68 uint8_t intp_use_natural[6];
69 uint8_t intp_reverse[6];
71 uint8_t pad[8];
72
73 // ~96 bytes. Comfortable margin under FZM_STAGE_CONFIG_SIZE (128 B);
74 // the static_assert below is the source of truth.
75
78 input_type(DataType::FLOAT32), code_type(DataType::UINT16),
79 ndim(3), eb_mode(0),
80 dim_x(0), dim_y(1), dim_z(1),
82 user_eb(0.0f), value_base(0.0f),
83 intp_alpha(1.75), intp_beta(4.0),
84 intp_use_md{1, 1, 0, 0, 0, 0},
85 intp_use_natural{0, 0, 0, 0, 0, 0},
86 intp_reverse{0, 0, 0, 0, 0, 0},
87 auto_tuning_mode(0), pad{} {}
88};
89static_assert(sizeof(GInterpConfig) <= FZM_STAGE_CONFIG_SIZE,
90 "GInterpConfig must fit in FZM_STAGE_CONFIG_SIZE");
91
205template <typename TInput = float, typename TCode = uint16_t>
206class GInterpStage : public Stage {
207public:
208 struct Config {
209 float error_bound = 1e-3f;
221 int quant_radius = 0;
222 float outlier_capacity = 0.10f;
226 std::array<size_t, 3> dims = {0, 0, 0};
230 float precomputed_value_base = 0.0f;
250 uint8_t auto_tuning_mode = 0;
251
255 double manual_alpha = 0.0;
257 double manual_beta = 0.0;
258
259 Config() = default;
260 };
261
262 explicit GInterpStage(const Config& cfg = Config()) : config_(cfg) {
263 actual_output_sizes_.resize(4, 0);
264 }
265 ~GInterpStage() override;
266
267 // ── Stage interface ──────────────────────────────────────────────────────
269 fz::stream_t stream,
270 MemoryPool* pool,
271 const std::vector<void*>& inputs,
272 const std::vector<void*>& outputs,
273 const std::vector<size_t>& sizes
274 ) override;
275
276 void postStreamSync(fz::stream_t stream) override;
277
283 void onFinalize(size_t estimated_inlen, MemoryPool* pool) override;
284
285 size_t estimateDeviceFootprintBytes(size_t /*estimated_inlen*/) const override {
286 return needsProfilingScratch() ? kProfilingErrCount * sizeof(float) : 0;
287 }
288 size_t estimatePinnedFootprintBytes(size_t /*estimated_inlen*/) const override {
289 return needsProfilingScratch() ? kProfilingErrCount * sizeof(float) : 0;
290 }
291
292 std::string getName() const override { return "GInterp"; }
293 size_t getNumInputs() const override { return is_inverse_ ? 4 : 1; }
294 size_t getNumOutputs() const override { return is_inverse_ ? 1 : 4; }
295
296 std::vector<std::string> getOutputNames() const override {
297 return {"codes", "anchor", "outlier_vals", "outlier_idxs"};
298 }
299
300 std::vector<size_t> estimateOutputSizes(
301 const std::vector<size_t>& input_sizes
302 ) const override;
303
304 std::unordered_map<std::string, size_t> getActualOutputSizesByName() const override {
305 auto names = getOutputNames();
306 std::unordered_map<std::string, size_t> r;
307 for (size_t i = 0; i < names.size() && i < actual_output_sizes_.size(); i++)
308 r[names[i]] = actual_output_sizes_[i];
309 return r;
310 }
311 size_t getActualOutputSize(int index) const override {
312 return (index >= 0 && index < static_cast<int>(actual_output_sizes_.size()))
313 ? actual_output_sizes_[index] : 0;
314 }
315
316 void saveState() override { saved_output_sizes_ = actual_output_sizes_; }
317 void restoreState() override { actual_output_sizes_ = saved_output_sizes_; }
318
319 // ── Setters ──────────────────────────────────────────────────────────────
320 void setErrorBound(float eb) { config_.error_bound = eb; }
321 void setQuantRadius(int radius) { config_.quant_radius = radius; }
322 void setOutlierCapacity(float cap) { config_.outlier_capacity = cap; }
328 config_.eb_mode = resolveApproxRelMode(m, "GInterpStage");
329 }
330 void setValueBase(float v) { config_.precomputed_value_base = v; }
335 void setAutoTuning(uint8_t mode) { config_.auto_tuning_mode = mode; }
340 void setManualAlphaBeta(double alpha, double beta) {
341 config_.manual_alpha = alpha;
342 config_.manual_beta = beta;
343 }
344 void setDims(const std::array<size_t, 3>& dims) override;
345 void setDims(size_t x, size_t y, size_t z) {
346 setDims(std::array<size_t, 3>{x, y, z});
347 }
348
349 float getErrorBound() const { return config_.error_bound; }
350 int getQuantRadius() const { return config_.quant_radius; }
351 float getOutlierCapacity() const { return config_.outlier_capacity; }
352 ErrorBoundMode getErrorBoundMode() const { return config_.eb_mode; }
353 float getValueBase() const { return config_.precomputed_value_base; }
354 uint8_t getAutoTuningMode() const { return config_.auto_tuning_mode; }
355 std::array<size_t, 3> getDims() const { return config_.dims; }
356
357 void setInverse(bool inv) override { is_inverse_ = inv; }
358 bool isInverse() const override { return is_inverse_; }
359
363 int ndim() const {
364 if (config_.dims[2] > 1) return 3;
365 if (config_.dims[1] > 1) return 2;
366 return 1;
367 }
368
369 // ── Type / Serialization ─────────────────────────────────────────────────
370 uint16_t getStageTypeId() const override {
371 return static_cast<uint16_t>(StageType::G_INTERP);
372 }
373
374 uint8_t getOutputDataType(size_t output_index) const override {
375 switch (output_index) {
376 case 0: return static_cast<uint8_t>(codeDataType()); // codes
377 case 1: return static_cast<uint8_t>(inputDataType()); // anchor
378 case 2: return static_cast<uint8_t>(inputDataType()); // outlier_vals
379 case 3: return static_cast<uint8_t>(DataType::UINT32); // outlier_idxs
380 default: return static_cast<uint8_t>(DataType::UINT8);
381 }
382 }
383 uint8_t getInputDataType(size_t /*input_index*/) const override {
384 return static_cast<uint8_t>(inputDataType());
385 }
386
387 size_t serializeHeader(size_t output_index, uint8_t* buf, size_t max_size) const override;
388 void deserializeHeader(const uint8_t* buf, size_t size) override;
389 size_t getMaxHeaderSize(size_t /*output_index*/) const override {
390 return sizeof(GInterpConfig);
391 }
392
409 bool isGraphCompatible() const override {
410 const bool tune_ok = (config_.auto_tuning_mode == 0 ||
411 config_.auto_tuning_mode == 5);
412 const bool radius_ok = config_.quant_radius > 0;
413 const bool eb_ok = (config_.eb_mode == ErrorBoundMode::ABS) ||
414 (config_.precomputed_value_base > 0.0f);
415 const bool mode5_alpha_ok = (config_.auto_tuning_mode != 5) ||
416 (config_.manual_alpha > 0.0);
417 return tune_ok && radius_ok && eb_ok && mode5_alpha_ok;
418 }
419
420private:
421 Config config_;
422 std::vector<size_t> actual_output_sizes_;
423 std::vector<size_t> saved_output_sizes_;
424
425 bool is_inverse_ = false;
426 size_t num_elements_ = 0;
427 uint32_t actual_outlier_count_ = 0;
434 uint32_t* d_outlier_count_scratch_ = nullptr;
435
438 TInput computed_abs_eb_ = 0;
441 float computed_value_base_ = 0.0f;
442
444 std::array<size_t, 3> anchor_dims_ = {0, 0, 0};
445
446 // ── Phase 2: auto-tuning state ──────────────────────────────────────────
449 static constexpr size_t kProfilingErrCount = 36;
450 float* d_profiling_errors_ = nullptr;
451 float* h_profiling_errors_ = nullptr;
455 MemoryPool* persistent_pool_ = nullptr;
460 std::weak_ptr<const void> persistent_pool_alive_;
461
462
467 double resolved_alpha_ = 1.75;
468 double resolved_beta_ = 4.0;
469 uint8_t resolved_use_md_[6] = {1, 1, 0, 0, 0, 0};
470 uint8_t resolved_use_natural_[6] = {0, 0, 0, 0, 0, 0};
471 uint8_t resolved_reverse_[6] = {0, 0, 0, 0, 0, 0};
472
475 bool needsProfilingScratch() const {
476 const uint8_t m = config_.auto_tuning_mode;
477 return m == 1 || m == 2 || m == 3 || m == 4;
478 }
479
482 void initProfilingScratch(MemoryPool* pool);
484 void initOutlierCountScratch(MemoryPool* pool);
488 INTERPOLATION_PARAMS buildIntpParam() const;
489
490 static DataType inputDataType() {
491 if (std::is_same<TInput, float>::value) return DataType::FLOAT32;
492 if (std::is_same<TInput, double>::value) return DataType::FLOAT64;
493 return DataType::FLOAT32;
494 }
495 static DataType codeDataType() {
496 if (std::is_same<TCode, uint8_t>::value) return DataType::UINT8;
497 if (std::is_same<TCode, uint16_t>::value) return DataType::UINT16;
498 if (std::is_same<TCode, uint32_t>::value) return DataType::UINT32;
499 return DataType::UINT16;
500 }
501 size_t getMaxOutlierCount(size_t n) const {
502 return static_cast<size_t>(std::ceil(n * config_.outlier_capacity));
503 }
504};
505
506extern template class GInterpStage<float, uint8_t>;
507extern template class GInterpStage<float, uint16_t>;
508extern template class GInterpStage<float, uint32_t>;
509extern template class GInterpStage<double, uint8_t>;
510extern template class GInterpStage<double, uint16_t>;
511extern template class GInterpStage<double, uint32_t>;
512
513} // namespace fz
Definition ginterp_stage.h:206
void postStreamSync(fz::stream_t stream) override
size_t serializeHeader(size_t output_index, uint8_t *buf, size_t max_size) const override
size_t getActualOutputSize(int index) const override
Definition ginterp_stage.h:311
void setDims(const std::array< size_t, 3 > &dims) override
uint8_t getInputDataType(size_t) const override
Definition ginterp_stage.h:383
uint8_t getOutputDataType(size_t output_index) const override
Definition ginterp_stage.h:374
std::unordered_map< std::string, size_t > getActualOutputSizesByName() const override
Definition ginterp_stage.h:304
size_t getMaxHeaderSize(size_t) const override
Definition ginterp_stage.h:389
void setAutoTuning(uint8_t mode)
Definition ginterp_stage.h:335
void onFinalize(size_t estimated_inlen, MemoryPool *pool) override
bool isGraphCompatible() const override
Definition ginterp_stage.h:409
void setManualAlphaBeta(double alpha, double beta)
Definition ginterp_stage.h:340
size_t estimateDeviceFootprintBytes(size_t) const override
Definition ginterp_stage.h:285
std::vector< std::string > getOutputNames() const override
Definition ginterp_stage.h:296
void saveState() override
Definition ginterp_stage.h:316
uint16_t getStageTypeId() const override
Definition ginterp_stage.h:370
int ndim() const
Definition ginterp_stage.h:363
void setErrorBoundMode(ErrorBoundMode m)
Definition ginterp_stage.h:327
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 ginterp_stage.h:357
std::vector< size_t > estimateOutputSizes(const std::vector< size_t > &input_sizes) const override
void deserializeHeader(const uint8_t *buf, size_t size) override
size_t estimatePinnedFootprintBytes(size_t) const override
Definition ginterp_stage.h:288
std::string getName() const override
Definition ginterp_stage.h:292
Definition mempool.h:82
Definition stage.h:30
FZM binary file format definitions — structs, enums, and helpers.
Fused Lorenzo predictor and quantizer stage.
Definition algorithms.h:48
ErrorBoundMode
Definition lorenzo_quant.h:40
@ ABS
Absolute error bound.
ErrorBoundMode resolveApproxRelMode(ErrorBoundMode mode, const char *stage_name)
Definition lorenzo_quant.h:61
constexpr size_t FZM_STAGE_CONFIG_SIZE
Per-stage serialized config slot (bytes)
Definition fzm_format.h:65
@ G_INTERP
Spline interpolation predictor + quantizer (cuSZ-Hi G-Interp)
DataType
Element data type identifiers used in buffer and stage descriptors.
Definition fzm_format.h:117
Base class interface for all compression stages.
Definition cusz_type_subset.h:32
Definition ginterp_stage.h:40
float user_eb
Original user-specified bound (before mode conversion).
Definition ginterp_stage.h:57
uint8_t eb_mode
ErrorBoundMode cast to uint8_t.
Definition ginterp_stage.h:50
uint8_t auto_tuning_mode
0=off, 1=cheap, 3=full, 4=full+alpha sweep, 5+=manual α/β.
Definition ginterp_stage.h:70
float value_base
value_range (NOA) / max(|data|) (REL) used in conversion.
Definition ginterp_stage.h:58
uint8_t intp_use_md[6]
Resolved use_md[level], booleans as u8.
Definition ginterp_stage.h:67
uint8_t pad[8]
Reserved for future fields (alignment also).
Definition ginterp_stage.h:71
uint32_t dim_x
X (fast) dimension.
Definition ginterp_stage.h:51
uint32_t anchor_dim_x
Anchor grid X extent.
Definition ginterp_stage.h:54
double error_bound
Definition ginterp_stage.h:42
uint32_t dim_z
Z dimension.
Definition ginterp_stage.h:53
double intp_alpha
Resolved alpha (auto-tuned from rel_eb or fixed).
Definition ginterp_stage.h:65
uint32_t num_elements
Total element count (= dim_x*dim_y*dim_z).
Definition ginterp_stage.h:45
uint32_t anchor_dim_y
Anchor grid Y extent.
Definition ginterp_stage.h:55
uint32_t outlier_count
Actual outlier count (post-execute).
Definition ginterp_stage.h:46
uint32_t anchor_dim_z
Anchor grid Z extent.
Definition ginterp_stage.h:56
double intp_beta
Resolved beta (default 4.0).
Definition ginterp_stage.h:66
uint32_t quant_radius
Quantization radius (codes lie in [0, 2*radius)).
Definition ginterp_stage.h:44
DataType code_type
Quant code type (1 B).
Definition ginterp_stage.h:48
uint8_t ndim
Spatial dimensionality (3 in MVP).
Definition ginterp_stage.h:49
uint32_t dim_y
Y dimension.
Definition ginterp_stage.h:52
DataType input_type
Float input type (1 B).
Definition ginterp_stage.h:47
Backend-neutral GPU type aliases.