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 "backend/types.h"
31#include <cstdint>
32#include <cstring>
33#include <memory>
34#include <stdexcept>
35#include <string>
36#include <unordered_map>
37#include <vector>
38
39namespace fz {
40
56class RZEStage : public Stage {
57public:
58 RZEStage()
59 : is_inverse_(false)
60 , chunk_size_(16384)
61 , word_size_(1)
62 , actual_output_size_(0)
63 , cached_orig_bytes_(0)
64 , d_scratch_(nullptr)
65 , d_sizes_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 // Variable-length coder = the sink that terminates a chunk-cooperative fused
90 // chain. block_size is the chunk in bytes; any byte-word chunk_size the
91 // fusion harness supports fuses — see chunk_geometry.h's kSupportedChunkBytes.
92 FusionSpec getFusionSpec() const override {
93 if (is_inverse_ || word_size_ != 1 ||
94 (chunk_size_ != 4096u && chunk_size_ != 8192u && chunk_size_ != 16384u)) return {};
95 return FusionSpec{FusionAccess::SegmentCodec, chunk_size_};
96 }
97
99 FusedOpDecl getFusedOp() const override {
100 if (!getFusionSpec().fusable()) return {};
101 return FusedOpDecl{FusionStrategy::ChunkCooperative, "RZECoder",
102 "fused/chunk_fusion/chunk_fusion.cuh", {}};
103 }
105 void setFusedArchiveResult(size_t archive_bytes, size_t orig_bytes) override {
106 setFusedResult(archive_bytes, orig_bytes);
107 }
113 void setFusedResult(size_t archive_bytes, size_t orig_bytes) {
114 actual_output_size_ = archive_bytes;
115 cached_orig_bytes_ = static_cast<uint32_t>(orig_bytes);
116 tail_readback_pending_ = false;
117 }
118
119 // ── Execution ──────────────────────────────────────────────────────────
121 fz::stream_t stream,
122 MemoryPool* pool,
123 const std::vector<void*>& inputs,
124 const std::vector<void*>& outputs,
125 const std::vector<size_t>& sizes
126 ) override;
127 void postStreamSync(fz::stream_t stream) override;
128
129 // ── Metadata ───────────────────────────────────────────────────────────
130 std::string getName() const override { return "RZE"; }
131 size_t getNumInputs() const override { return 1; }
132 size_t getNumOutputs() const override { return 1; }
133
134 std::vector<size_t> estimateOutputSizes(
135 const std::vector<size_t>& input_sizes
136 ) const override {
137 if (is_inverse_) {
138 if (cached_orig_bytes_ > 0)
139 return {static_cast<size_t>(cached_orig_bytes_)};
140 return {input_sizes.empty() ? 0 : input_sizes[0]};
141 }
142 // Forward: worst case = original data + stream header.
143 const size_t n_bytes = input_sizes.empty() ? 0 : input_sizes[0];
144 const size_t n_chunks = (n_bytes + chunk_size_ - 1) / chunk_size_;
145 const size_t hdr = 4 + 4 + 4 * n_chunks;
146 // postStreamSync()/getActualOutputSizesByName() always round the final
147 // size up to a 4-byte boundary and zero-fill the pad, even when the
148 // real total isn't already aligned (e.g. a partial final chunk stored
149 // raw at a byte count that isn't a multiple of 4) -- reserve that pad
150 // here too, or the caller's allocation is up to 3 bytes short and the
151 // memset in postStreamSync writes out of bounds.
152 const size_t worst = n_bytes + hdr;
153 return {(worst + 3) & ~size_t(3)};
154 }
155
156 std::unordered_map<std::string, size_t>
158 size_t getActualOutputSize(int index) const override;
159
168 const std::vector<size_t>& input_sizes
169 ) const override {
170 if (is_inverse_ || input_sizes.empty()) return 0;
171 const size_t in_bytes = input_sizes[0];
172 const size_t n_chunks = (in_bytes + chunk_size_ - 1) / chunk_size_;
173 return n_chunks * (static_cast<size_t>(chunk_size_) + 3 * sizeof(uint32_t));
174 }
175
176 uint16_t getStageTypeId() const override {
177 return static_cast<uint16_t>(StageType::RZE);
178 }
179
180 uint8_t getOutputDataType(size_t) const override {
181 return static_cast<uint8_t>(DataType::UINT8);
182 }
183
184 // ── Serialization ──────────────────────────────────────────────────────
186 size_t output_index, uint8_t* buf, size_t max_size
187 ) const override {
188 (void)output_index;
189 if (max_size < 9) return 0;
190 std::memcpy(buf, &chunk_size_, sizeof(uint32_t));
191 buf[4] = word_size_;
192 std::memcpy(buf + 5, &cached_orig_bytes_, sizeof(uint32_t));
193 return 9;
194 }
195
196 void deserializeHeader(const uint8_t* buf, size_t size) override {
197 if (size >= 4) std::memcpy(&chunk_size_, buf, sizeof(uint32_t));
198 if (size >= 5) word_size_ = buf[4];
199 if (size >= 9) std::memcpy(&cached_orig_bytes_, buf + 5, sizeof(uint32_t));
200 }
201
202 size_t getMaxHeaderSize(size_t) const override { return 9; }
203
204 void saveState() override {
205 saved_chunk_size_ = chunk_size_;
206 saved_word_size_ = word_size_;
207 saved_cached_orig_bytes_ = cached_orig_bytes_;
208 }
209
210 void restoreState() override {
211 chunk_size_ = saved_chunk_size_;
212 word_size_ = saved_word_size_;
213 cached_orig_bytes_ = saved_cached_orig_bytes_;
214 }
215
216private:
217 bool is_inverse_;
218 uint32_t chunk_size_;
219 uint32_t saved_chunk_size_ = 0;
220 uint8_t word_size_;
221 uint8_t saved_word_size_ = 0;
222 size_t actual_output_size_;
223 uint32_t cached_orig_bytes_ = 0;
224 uint32_t saved_cached_orig_bytes_ = 0;
225
226 // ── Persistent forward scratch buffers ───────────────────────────────────
227 uint8_t* d_scratch_;
228 uint32_t* d_sizes_dev_;
229 uint32_t* d_dst_off_dev_;
230 mutable bool tail_readback_pending_ = false;
231 mutable fz::stream_t tail_readback_stream_ = nullptr;
232 mutable uint32_t tail_last_index_ = 0;
233 mutable uint32_t tail_header_size_ = 0;
234 mutable uint8_t* tail_output_ptr_ = nullptr;
235 size_t scratch_capacity_;
236 MemoryPool* scratch_pool_owner_ = nullptr;
237 bool scratch_from_pool_ = false;
239 std::weak_ptr<const void> scratch_alive_;
240};
241
242} // namespace fz
Definition mempool.h:82
Definition rze_stage.h:56
FusedOpDecl getFusedOp() const override
Chunk-cooperative coder op (the swappable variable-length sink). Stateless.
Definition rze_stage.h:99
FusionSpec getFusionSpec() const override
Definition rze_stage.h:92
uint8_t getOutputDataType(size_t) const override
Definition rze_stage.h:180
void deserializeHeader(const uint8_t *buf, size_t size) override
Definition rze_stage.h:196
bool isGraphCompatible() const override
Definition rze_stage.h:79
void setFusedArchiveResult(size_t archive_bytes, size_t orig_bytes) override
Base-class tail hook → the existing coder result setter (archive, orig).
Definition rze_stage.h:105
void setInverse(bool inv) override
Definition rze_stage.h:73
void setFusedResult(size_t archive_bytes, size_t orig_bytes)
Definition rze_stage.h:113
size_t getMaxHeaderSize(size_t) const override
Definition rze_stage.h:202
size_t serializeHeader(size_t output_index, uint8_t *buf, size_t max_size) const override
Definition rze_stage.h:185
uint16_t getStageTypeId() const override
Definition rze_stage.h:176
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:204
std::string getName() const override
Definition rze_stage.h:130
std::vector< size_t > estimateOutputSizes(const std::vector< size_t > &input_sizes) const override
Definition rze_stage.h:134
void postStreamSync(fz::stream_t stream) override
size_t getActualOutputSize(int index) const override
size_t estimateScratchBytes(const std::vector< size_t > &input_sizes) const override
Definition rze_stage.h:167
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
Definition stage.h:31
FZM binary file format definitions — structs, enums, and helpers.
Definition dag.h:24
Base class interface for all compression stages.
A stage's contribution to a generated fused kernel — the device-op it maps to, where its source lives...
Definition fusion.h:161
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.