FZGPUModules 2.0
GPU-accelerated modular compression pipelines
Loading...
Searching...
No Matches
rze_stage.h
Go to the documentation of this file.
1#pragma once
2
28#include "stage/stage.h"
29#include "fzm_format.h"
30#include <cuda_runtime.h>
31#include <cstdint>
32#include <cstring>
33#include <stdexcept>
34#include <string>
35#include <unordered_map>
36#include <vector>
37
38namespace fz {
39
55class RZEStage : public Stage {
56public:
57 RZEStage()
58 : is_inverse_(false)
59 , chunk_size_(16384)
60 , word_size_(1)
61 , actual_output_size_(0)
62 , cached_orig_bytes_(0)
63 , d_scratch_(nullptr)
64 , d_sizes_dev_(nullptr)
65 , d_clean_dev_(nullptr)
66 , d_dst_off_dev_(nullptr)
67 , scratch_capacity_(0)
68 {}
69
70 ~RZEStage() override;
71
72 // ── Stage control ──────────────────────────────────────────────────────
73 void setInverse(bool inv) override { is_inverse_ = inv; }
74 bool isInverse() const override { return is_inverse_; }
75
79 bool isGraphCompatible() const override { return !is_inverse_; }
80
81 void setChunkSize(size_t bytes) { chunk_size_ = static_cast<uint32_t>(bytes); }
82 void setWordSize(size_t bytes) { word_size_ = static_cast<uint8_t>(bytes); }
83
84 size_t getChunkSize() const { return chunk_size_; }
85 size_t getRequiredInputAlignment() const override { return chunk_size_; }
86 int getWordSize() const { return static_cast<int>(word_size_); }
87 uint32_t getCachedOrigBytes() const { return cached_orig_bytes_; }
88
89 // ── Execution ──────────────────────────────────────────────────────────
90 void execute(
91 cudaStream_t stream,
92 MemoryPool* pool,
93 const std::vector<void*>& inputs,
94 const std::vector<void*>& outputs,
95 const std::vector<size_t>& sizes
96 ) override;
97 void postStreamSync(cudaStream_t stream) override;
98
99 // ── Metadata ───────────────────────────────────────────────────────────
100 std::string getName() const override { return "RZE"; }
101 size_t getNumInputs() const override { return 1; }
102 size_t getNumOutputs() const override { return 1; }
103
104 std::vector<size_t> estimateOutputSizes(
105 const std::vector<size_t>& input_sizes
106 ) const override {
107 if (is_inverse_) {
108 if (cached_orig_bytes_ > 0)
109 return {static_cast<size_t>(cached_orig_bytes_)};
110 return {input_sizes.empty() ? 0 : input_sizes[0]};
111 }
112 // Forward: worst case = original data + stream header.
113 const size_t n_bytes = input_sizes.empty() ? 0 : input_sizes[0];
114 const size_t n_chunks = (n_bytes + chunk_size_ - 1) / chunk_size_;
115 const size_t hdr = 4 + 4 + 4 * n_chunks;
116 return {n_bytes + hdr};
117 }
118
119 std::unordered_map<std::string, size_t>
121 size_t getActualOutputSize(int index) const override;
122
132 const std::vector<size_t>& input_sizes
133 ) const override {
134 if (is_inverse_ || input_sizes.empty()) return 0;
135 const size_t in_bytes = input_sizes[0];
136 const size_t n_chunks = (in_bytes + chunk_size_ - 1) / chunk_size_;
137 return n_chunks * (static_cast<size_t>(chunk_size_) + 3 * sizeof(uint32_t));
138 }
139
140 uint16_t getStageTypeId() const override {
141 return static_cast<uint16_t>(StageType::RZE);
142 }
143
144 uint8_t getOutputDataType(size_t) const override {
145 return static_cast<uint8_t>(DataType::UINT8);
146 }
147
148 // ── Serialization ──────────────────────────────────────────────────────
150 size_t output_index, uint8_t* buf, size_t max_size
151 ) const override {
152 (void)output_index;
153 if (max_size < 9) return 0;
154 std::memcpy(buf, &chunk_size_, sizeof(uint32_t));
155 buf[4] = word_size_;
156 std::memcpy(buf + 5, &cached_orig_bytes_, sizeof(uint32_t));
157 return 9;
158 }
159
160 void deserializeHeader(const uint8_t* buf, size_t size) override {
161 if (size >= 4) std::memcpy(&chunk_size_, buf, sizeof(uint32_t));
162 if (size >= 5) word_size_ = buf[4];
163 if (size >= 9) std::memcpy(&cached_orig_bytes_, buf + 5, sizeof(uint32_t));
164 }
165
166 size_t getMaxHeaderSize(size_t) const override { return 9; }
167
168 void saveState() override {
169 saved_chunk_size_ = chunk_size_;
170 saved_word_size_ = word_size_;
171 saved_cached_orig_bytes_ = cached_orig_bytes_;
172 }
173
174 void restoreState() override {
175 chunk_size_ = saved_chunk_size_;
176 word_size_ = saved_word_size_;
177 cached_orig_bytes_ = saved_cached_orig_bytes_;
178 }
179
180private:
181 bool is_inverse_;
182 uint32_t chunk_size_;
183 uint32_t saved_chunk_size_ = 0;
184 uint8_t word_size_;
185 uint8_t saved_word_size_ = 0;
186 size_t actual_output_size_;
187 uint32_t cached_orig_bytes_ = 0;
188 uint32_t saved_cached_orig_bytes_ = 0;
189
190 // ── Persistent forward scratch buffers ───────────────────────────────────
191 uint8_t* d_scratch_;
192 uint32_t* d_sizes_dev_;
193 uint32_t* d_clean_dev_;
194 uint32_t* d_dst_off_dev_;
195 mutable bool tail_readback_pending_ = false;
196 mutable cudaStream_t tail_readback_stream_ = nullptr;
197 mutable uint32_t tail_last_index_ = 0;
198 mutable uint8_t* tail_output_ptr_ = nullptr;
199 size_t scratch_capacity_;
200 MemoryPool* scratch_pool_owner_ = nullptr;
201 bool scratch_from_pool_ = false;
202};
203
204} // namespace fz
Definition mempool.h:82
Definition rze_stage.h:55
void execute(cudaStream_t stream, MemoryPool *pool, const std::vector< void * > &inputs, const std::vector< void * > &outputs, const std::vector< size_t > &sizes) override
uint8_t getOutputDataType(size_t) const override
Definition rze_stage.h:144
void deserializeHeader(const uint8_t *buf, size_t size) override
Definition rze_stage.h:160
bool isGraphCompatible() const override
Definition rze_stage.h:79
void setInverse(bool inv) override
Definition rze_stage.h:73
size_t getMaxHeaderSize(size_t) const override
Definition rze_stage.h:166
size_t serializeHeader(size_t output_index, uint8_t *buf, size_t max_size) const override
Definition rze_stage.h:149
uint16_t getStageTypeId() const override
Definition rze_stage.h:140
std::unordered_map< std::string, size_t > getActualOutputSizesByName() const override
size_t getRequiredInputAlignment() const override
Definition rze_stage.h:85
void saveState() override
Definition rze_stage.h:168
std::string getName() const override
Definition rze_stage.h:100
std::vector< size_t > estimateOutputSizes(const std::vector< size_t > &input_sizes) const override
Definition rze_stage.h:104
size_t getActualOutputSize(int index) const override
size_t estimateScratchBytes(const std::vector< size_t > &input_sizes) const override
Definition rze_stage.h:131
void postStreamSync(cudaStream_t stream) override
Definition stage.h:30
FZM binary file format definitions — structs, enums, and helpers.
Definition fzm_format.h:25
Base class interface for all compression stages.