FZGPUModules 2.0
GPU-accelerated modular compression pipelines
Loading...
Searching...
No Matches
rare_stage.h
Go to the documentation of this file.
1#pragma once
2
34#include "stage/stage.h"
35#include "fzm_format.h"
36#include "backend/types.h"
37#include <cstdint>
38#include <cstring>
39#include <memory>
40#include <stdexcept>
41#include <string>
42#include <unordered_map>
43#include <vector>
44
45namespace fz {
46
64class RAREStage : public Stage {
65public:
66 RAREStage()
67 : is_inverse_(false)
68 , chunk_size_(16384)
69 , word_size_(1)
70 , actual_output_size_(0)
71 , cached_orig_bytes_(0)
72 , d_scratch_(nullptr)
73 , d_sizes_dev_(nullptr)
74 , d_clean_dev_(nullptr)
75 , d_dst_off_dev_(nullptr)
76 , scratch_capacity_(0)
77 {}
78
79 ~RAREStage() override;
80
81 // ── Stage control ──────────────────────────────────────────────────────
82 void setInverse(bool inv) override { is_inverse_ = inv; }
83 bool isInverse() const override { return is_inverse_; }
84
88 bool isGraphCompatible() const override { return !is_inverse_; }
89
90 void setChunkSize(size_t bytes) { chunk_size_ = static_cast<uint32_t>(bytes); }
91 void setWordSize(size_t bytes) { word_size_ = static_cast<uint8_t>(bytes); }
92
93 size_t getChunkSize() const { return chunk_size_; }
94 size_t getRequiredInputAlignment() const override { return chunk_size_; }
95 int getWordSize() const { return static_cast<int>(word_size_); }
96 uint32_t getCachedOrigBytes() const { return cached_orig_bytes_; }
97
98 // Chunk-cooperative variable-length coder (the swappable sink) — identical
99 // machinery to RZE/RRE. Any byte-word chunk_size the fusion harness supports
100 // fuses (matches the fused RARECoder<ChunkBytes> device op); see
101 // chunk_geometry.h's kSupportedChunkBytes.
102 FusionSpec getFusionSpec() const override {
103 if (is_inverse_ || word_size_ != 1 ||
104 (chunk_size_ != 4096u && chunk_size_ != 8192u && chunk_size_ != 16384u)) return {};
105 return FusionSpec{FusionAccess::SegmentCodec, chunk_size_};
106 }
107 FusedOpDecl getFusedOp() const override {
108 if (!getFusionSpec().fusable()) return {};
109 return FusedOpDecl{FusionStrategy::ChunkCooperative, "RARECoder",
110 "fused/chunk_fusion/chunk_fusion.cuh", {}};
111 }
114 void setFusedResult(size_t archive_bytes, size_t orig_bytes) {
115 actual_output_size_ = archive_bytes;
116 cached_orig_bytes_ = static_cast<uint32_t>(orig_bytes);
117 tail_readback_pending_ = false;
118 }
119 void setFusedArchiveResult(size_t archive_bytes, size_t orig_bytes) override {
120 setFusedResult(archive_bytes, orig_bytes);
121 }
122
123 // ── Execution ──────────────────────────────────────────────────────────
125 cudaStream_t stream,
126 MemoryPool* pool,
127 const std::vector<void*>& inputs,
128 const std::vector<void*>& outputs,
129 const std::vector<size_t>& sizes
130 ) override;
131 void postStreamSync(cudaStream_t stream) override;
132
133 // ── Metadata ───────────────────────────────────────────────────────────
134 std::string getName() const override { return "RARE"; }
135 size_t getNumInputs() const override { return 1; }
136 size_t getNumOutputs() const override { return 1; }
137
138 std::vector<size_t> estimateOutputSizes(
139 const std::vector<size_t>& input_sizes
140 ) const override {
141 if (is_inverse_) {
142 if (cached_orig_bytes_ > 0)
143 return {static_cast<size_t>(cached_orig_bytes_)};
144 return {input_sizes.empty() ? 0 : input_sizes[0]};
145 }
146 // Forward: worst case = original data + stream header.
147 const size_t n_bytes = input_sizes.empty() ? 0 : input_sizes[0];
148 const size_t n_chunks = (n_bytes + chunk_size_ - 1) / chunk_size_;
149 const size_t hdr = 4 + 4 + 4 * n_chunks;
150 // postStreamSync()/getActualOutputSizesByName() always round the final
151 // size up to a 4-byte boundary and zero-fill the pad, even when the
152 // real total isn't already aligned (e.g. a partial final chunk stored
153 // raw at a byte count that isn't a multiple of 4) -- reserve that pad
154 // here too, or the caller's allocation is up to 3 bytes short and the
155 // memset in postStreamSync writes out of bounds.
156 const size_t worst = n_bytes + hdr;
157 return {(worst + 3) & ~size_t(3)};
158 }
159
160 std::unordered_map<std::string, size_t>
162 size_t getActualOutputSize(int index) const override;
163
173 const std::vector<size_t>& input_sizes
174 ) const override {
175 if (is_inverse_ || input_sizes.empty()) return 0;
176 const size_t in_bytes = input_sizes[0];
177 const size_t n_chunks = (in_bytes + chunk_size_ - 1) / chunk_size_;
178 return n_chunks * (static_cast<size_t>(chunk_size_) + 3 * sizeof(uint32_t));
179 }
180
181 uint16_t getStageTypeId() const override {
182 return static_cast<uint16_t>(StageType::RARE);
183 }
184
185 uint8_t getOutputDataType(size_t) const override {
186 return static_cast<uint8_t>(DataType::UINT8);
187 }
188
189 // ── Serialization ──────────────────────────────────────────────────────
191 size_t output_index, uint8_t* buf, size_t max_size
192 ) const override {
193 (void)output_index;
194 if (max_size < 9) return 0;
195 std::memcpy(buf, &chunk_size_, sizeof(uint32_t));
196 buf[4] = word_size_;
197 std::memcpy(buf + 5, &cached_orig_bytes_, sizeof(uint32_t));
198 return 9;
199 }
200
201 void deserializeHeader(const uint8_t* buf, size_t size) override {
202 if (size >= 4) std::memcpy(&chunk_size_, buf, sizeof(uint32_t));
203 if (size >= 5) word_size_ = buf[4];
204 if (size >= 9) std::memcpy(&cached_orig_bytes_, buf + 5, sizeof(uint32_t));
205 }
206
207 size_t getMaxHeaderSize(size_t) const override { return 9; }
208
209 void saveState() override {
210 saved_chunk_size_ = chunk_size_;
211 saved_word_size_ = word_size_;
212 saved_cached_orig_bytes_ = cached_orig_bytes_;
213 }
214
215 void restoreState() override {
216 chunk_size_ = saved_chunk_size_;
217 word_size_ = saved_word_size_;
218 cached_orig_bytes_ = saved_cached_orig_bytes_;
219 }
220
221private:
222 bool is_inverse_;
223 uint32_t chunk_size_;
224 uint32_t saved_chunk_size_ = 0;
225 uint8_t word_size_;
226 uint8_t saved_word_size_ = 0;
227 size_t actual_output_size_;
228 uint32_t cached_orig_bytes_ = 0;
229 uint32_t saved_cached_orig_bytes_ = 0;
230
231 // ── Persistent forward scratch buffers ───────────────────────────────────
232 uint8_t* d_scratch_;
233 uint32_t* d_sizes_dev_;
234 uint32_t* d_clean_dev_;
235 uint32_t* d_dst_off_dev_;
236 mutable bool tail_readback_pending_ = false;
237 mutable cudaStream_t tail_readback_stream_ = nullptr;
238 mutable uint32_t tail_last_index_ = 0;
239 mutable uint8_t* tail_output_ptr_ = nullptr;
240 size_t scratch_capacity_;
241 MemoryPool* scratch_pool_owner_ = nullptr;
242 bool scratch_from_pool_ = false;
244 std::weak_ptr<const void> scratch_alive_;
245};
246
247} // namespace fz
Definition mempool.h:82
Definition rare_stage.h:64
void setFusedArchiveResult(size_t archive_bytes, size_t orig_bytes) override
Definition rare_stage.h:119
void setInverse(bool inv) override
Definition rare_stage.h:82
size_t getMaxHeaderSize(size_t) const override
Definition rare_stage.h:207
void saveState() override
Definition rare_stage.h:209
bool isGraphCompatible() const override
Definition rare_stage.h:88
void deserializeHeader(const uint8_t *buf, size_t size) override
Definition rare_stage.h:201
void postStreamSync(cudaStream_t stream) override
std::unordered_map< std::string, size_t > getActualOutputSizesByName() const override
uint8_t getOutputDataType(size_t) const override
Definition rare_stage.h:185
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 rare_stage.h:94
FusedOpDecl getFusedOp() const override
Definition rare_stage.h:107
size_t serializeHeader(size_t output_index, uint8_t *buf, size_t max_size) const override
Definition rare_stage.h:190
FusionSpec getFusionSpec() const override
Definition rare_stage.h:102
std::vector< size_t > estimateOutputSizes(const std::vector< size_t > &input_sizes) const override
Definition rare_stage.h:138
size_t estimateScratchBytes(const std::vector< size_t > &input_sizes) const override
Definition rare_stage.h:172
uint16_t getStageTypeId() const override
Definition rare_stage.h:181
size_t getActualOutputSize(int index) const override
std::string getName() const override
Definition rare_stage.h:134
void setFusedResult(size_t archive_bytes, size_t orig_bytes)
Definition rare_stage.h:114
Definition stage.h:31
FZM binary file format definitions — structs, enums, and helpers.
Definition dag.h:24
@ RARE
Adaptive top-bit matching generalization of RRE (LC component)
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.