FZGPUModules 2.0
GPU-accelerated modular compression pipelines
Loading...
Searching...
No Matches
raze_stage.h
Go to the documentation of this file.
1#pragma once
2
33#include "stage/stage.h"
34#include "fzm_format.h"
35#include "backend/types.h"
36#include <cstdint>
37#include <cstring>
38#include <stdexcept>
39#include <string>
40#include <unordered_map>
41#include <vector>
42
43namespace fz {
44
62class RAZEStage : public Stage {
63public:
64 RAZEStage()
65 : is_inverse_(false)
66 , chunk_size_(16384)
67 , word_size_(1)
68 , actual_output_size_(0)
69 , cached_orig_bytes_(0)
70 , d_scratch_(nullptr)
71 , d_sizes_dev_(nullptr)
72 , d_clean_dev_(nullptr)
73 , d_dst_off_dev_(nullptr)
74 , scratch_capacity_(0)
75 {}
76
77 ~RAZEStage() override;
78
79 // ── Stage control ──────────────────────────────────────────────────────
80 void setInverse(bool inv) override { is_inverse_ = inv; }
81 bool isInverse() const override { return is_inverse_; }
82
86 bool isGraphCompatible() const override { return !is_inverse_; }
87
88 void setChunkSize(size_t bytes) { chunk_size_ = static_cast<uint32_t>(bytes); }
89 void setWordSize(size_t bytes) { word_size_ = static_cast<uint8_t>(bytes); }
90
91 size_t getChunkSize() const { return chunk_size_; }
92 size_t getRequiredInputAlignment() const override { return chunk_size_; }
93 int getWordSize() const { return static_cast<int>(word_size_); }
94 uint32_t getCachedOrigBytes() const { return cached_orig_bytes_; }
95
96 // ── Execution ──────────────────────────────────────────────────────────
97 void execute(
98 cudaStream_t stream,
99 MemoryPool* pool,
100 const std::vector<void*>& inputs,
101 const std::vector<void*>& outputs,
102 const std::vector<size_t>& sizes
103 ) override;
104 void postStreamSync(cudaStream_t stream) override;
105
106 // ── Metadata ───────────────────────────────────────────────────────────
107 std::string getName() const override { return "RAZE"; }
108 size_t getNumInputs() const override { return 1; }
109 size_t getNumOutputs() const override { return 1; }
110
111 std::vector<size_t> estimateOutputSizes(
112 const std::vector<size_t>& input_sizes
113 ) const override {
114 if (is_inverse_) {
115 if (cached_orig_bytes_ > 0)
116 return {static_cast<size_t>(cached_orig_bytes_)};
117 return {input_sizes.empty() ? 0 : input_sizes[0]};
118 }
119 // Forward: worst case = original data + stream header.
120 const size_t n_bytes = input_sizes.empty() ? 0 : input_sizes[0];
121 const size_t n_chunks = (n_bytes + chunk_size_ - 1) / chunk_size_;
122 const size_t hdr = 4 + 4 + 4 * n_chunks;
123 // postStreamSync()/getActualOutputSizesByName() always round the final
124 // size up to a 4-byte boundary and zero-fill the pad, even when the
125 // real total isn't already aligned (e.g. a partial final chunk stored
126 // raw at a byte count that isn't a multiple of 4) -- reserve that pad
127 // here too, or the caller's allocation is up to 3 bytes short and the
128 // memset in postStreamSync writes out of bounds.
129 const size_t worst = n_bytes + hdr;
130 return {(worst + 3) & ~size_t(3)};
131 }
132
133 std::unordered_map<std::string, size_t>
135 size_t getActualOutputSize(int index) const override;
136
146 const std::vector<size_t>& input_sizes
147 ) const override {
148 if (is_inverse_ || input_sizes.empty()) return 0;
149 const size_t in_bytes = input_sizes[0];
150 const size_t n_chunks = (in_bytes + chunk_size_ - 1) / chunk_size_;
151 return n_chunks * (static_cast<size_t>(chunk_size_) + 3 * sizeof(uint32_t));
152 }
153
154 uint16_t getStageTypeId() const override {
155 return static_cast<uint16_t>(StageType::RAZE);
156 }
157
158 uint8_t getOutputDataType(size_t) const override {
159 return static_cast<uint8_t>(DataType::UINT8);
160 }
161
162 // ── Serialization ──────────────────────────────────────────────────────
164 size_t output_index, uint8_t* buf, size_t max_size
165 ) const override {
166 (void)output_index;
167 if (max_size < 9) return 0;
168 std::memcpy(buf, &chunk_size_, sizeof(uint32_t));
169 buf[4] = word_size_;
170 std::memcpy(buf + 5, &cached_orig_bytes_, sizeof(uint32_t));
171 return 9;
172 }
173
174 void deserializeHeader(const uint8_t* buf, size_t size) override {
175 if (size >= 4) std::memcpy(&chunk_size_, buf, sizeof(uint32_t));
176 if (size >= 5) word_size_ = buf[4];
177 if (size >= 9) std::memcpy(&cached_orig_bytes_, buf + 5, sizeof(uint32_t));
178 }
179
180 size_t getMaxHeaderSize(size_t) const override { return 9; }
181
182 void saveState() override {
183 saved_chunk_size_ = chunk_size_;
184 saved_word_size_ = word_size_;
185 saved_cached_orig_bytes_ = cached_orig_bytes_;
186 }
187
188 void restoreState() override {
189 chunk_size_ = saved_chunk_size_;
190 word_size_ = saved_word_size_;
191 cached_orig_bytes_ = saved_cached_orig_bytes_;
192 }
193
194private:
195 bool is_inverse_;
196 uint32_t chunk_size_;
197 uint32_t saved_chunk_size_ = 0;
198 uint8_t word_size_;
199 uint8_t saved_word_size_ = 0;
200 size_t actual_output_size_;
201 uint32_t cached_orig_bytes_ = 0;
202 uint32_t saved_cached_orig_bytes_ = 0;
203
204 // ── Persistent forward scratch buffers ───────────────────────────────────
205 uint8_t* d_scratch_;
206 uint32_t* d_sizes_dev_;
207 uint32_t* d_clean_dev_;
208 uint32_t* d_dst_off_dev_;
209 mutable bool tail_readback_pending_ = false;
210 mutable cudaStream_t tail_readback_stream_ = nullptr;
211 mutable uint32_t tail_last_index_ = 0;
212 mutable uint8_t* tail_output_ptr_ = nullptr;
213 size_t scratch_capacity_;
214 MemoryPool* scratch_pool_owner_ = nullptr;
215 bool scratch_from_pool_ = false;
216};
217
218} // namespace fz
Definition mempool.h:82
Definition raze_stage.h:62
std::vector< size_t > estimateOutputSizes(const std::vector< size_t > &input_sizes) const override
Definition raze_stage.h:111
std::unordered_map< std::string, size_t > getActualOutputSizesByName() const override
size_t getActualOutputSize(int index) const override
void execute(cudaStream_t stream, MemoryPool *pool, const std::vector< void * > &inputs, const std::vector< void * > &outputs, const std::vector< size_t > &sizes) override
size_t getRequiredInputAlignment() const override
Definition raze_stage.h:92
void setInverse(bool inv) override
Definition raze_stage.h:80
void postStreamSync(cudaStream_t stream) override
uint16_t getStageTypeId() const override
Definition raze_stage.h:154
void deserializeHeader(const uint8_t *buf, size_t size) override
Definition raze_stage.h:174
size_t getMaxHeaderSize(size_t) const override
Definition raze_stage.h:180
size_t serializeHeader(size_t output_index, uint8_t *buf, size_t max_size) const override
Definition raze_stage.h:163
std::string getName() const override
Definition raze_stage.h:107
bool isGraphCompatible() const override
Definition raze_stage.h:86
size_t estimateScratchBytes(const std::vector< size_t > &input_sizes) const override
Definition raze_stage.h:145
uint8_t getOutputDataType(size_t) const override
Definition raze_stage.h:158
void saveState() override
Definition raze_stage.h:182
Definition stage.h:30
FZM binary file format definitions — structs, enums, and helpers.
Definition algorithms.h:48
@ RAZE
Zero-Adaptive Reduction Encoding (LC framework, auto-k generalization of RZE)
Base class interface for all compression stages.
Backend-neutral GPU type aliases.