FZGPUModules 2.0
GPU-accelerated modular compression pipelines
Loading...
Searching...
No Matches
diff.h
Go to the documentation of this file.
1#pragma once
2
8#include "stage/stage.h"
9#include "fzm_format.h"
10#include "backend/types.h"
11#include <cstdint>
12#include <cstring>
13#include <type_traits>
14
15namespace fz {
16
28enum class FusionMode : uint8_t {
29 NEGABINARY = 0,
30 ZIGZAG = 1,
31};
32
70template<typename T = float, typename TOut = T, FusionMode Mode = FusionMode::NEGABINARY>
71class DifferenceStage : public Stage {
72 static_assert(
73 std::is_same_v<T, TOut> ||
74 (std::is_integral_v<T> && std::is_signed_v<T> &&
75 std::is_integral_v<TOut> && std::is_unsigned_v<TOut> &&
76 sizeof(T) == sizeof(TOut)),
77 "DifferenceStage: TOut must equal T, or T must be a signed integer "
78 "and TOut its unsigned counterpart of the same width (negabinary fusion).");
79public:
80 DifferenceStage() : actual_output_size_(0), is_inverse_(false), chunk_size_(0) {}
81
82 void setInverse(bool inverse) override { is_inverse_ = inverse; }
83 bool isInverse() const override { return is_inverse_; }
84
92 void setChunkSize(size_t bytes) { chunk_size_ = bytes; }
93 size_t getChunkSize() const { return chunk_size_; }
94 size_t getRequiredInputAlignment() const override {
95 return chunk_size_ > 0 ? chunk_size_ : 1;
96 }
97
98 // Region-local (chunk-cooperative) in two shapes:
99 // - T != TOut, Mode == NEGABINARY: the fused DiffNegabinary op (unchanged).
100 // - T == TOut (plain difference, no fused encode step): the fused DiffPlain
101 // op, gated to int32_t -- the chunk_fusion.cuh harness's shared buffers
102 // are a fixed uint32_t[4096] (16 KB / 4 B) shape, and DiffPlain's device
103 // op is written against that width (the NEGABINARY op has the same real
104 // constraint but happens to only be instantiated/used at int32 today, so
105 // it was never gated explicitly). Plain difference is the shape a coder
106 // that does its OWN final encode (e.g. GolombRiceCoder, which zigzags
107 // internally) needs -- a transform that also encoded would double-encode.
108 // ZIGZAG-fused (T != TOut, Mode == ZIGZAG) is deliberately NOT declared
109 // fusable: no chunk-cooperative coder currently consumes a pre-zigzagged
110 // transform output (GolombRiceCoder uses DiffPlain instead, precisely to
111 // avoid double-zigzagging itself), so there is nothing to test it against.
112 // block_size is the chunk in BYTES (the granularity the whole chunk chain
113 // shares).
114 FusionSpec getFusionSpec() const override {
115 if (is_inverse_ || chunk_size_ == 0) return {};
116 if constexpr (!std::is_same_v<T, TOut> && Mode == FusionMode::NEGABINARY)
117 return FusionSpec{FusionAccess::RegionLocal, static_cast<uint32_t>(chunk_size_)};
118 else if constexpr (std::is_same_v<T, TOut> && std::is_same_v<T, int32_t>)
119 return FusionSpec{FusionAccess::RegionLocal, static_cast<uint32_t>(chunk_size_)};
120 else
121 return {};
122 }
123
127 FusedOpDecl getFusedOp() const override {
128 if (!getFusionSpec().fusable()) return {};
129 const char* op_name = std::is_same_v<T, TOut> ? "DiffPlain" : "DiffNegabinary";
130 return FusedOpDecl{FusionStrategy::ChunkCooperative, op_name,
131 "fused/chunk_fusion/chunk_fusion.cuh", {}};
132 }
133
135 fz::stream_t stream,
136 MemoryPool* pool,
137 const std::vector<void*>& inputs,
138 const std::vector<void*>& outputs,
139 const std::vector<size_t>& sizes
140 ) override;
141
142
143 std::string getName() const override { return "Difference"; }
144 size_t getNumInputs() const override { return 1; }
145 size_t getNumOutputs() const override { return 1; }
146
147 std::vector<size_t> estimateOutputSizes(
148 const std::vector<size_t>& input_sizes
149 ) const override {
150 return {input_sizes[0]}; // size-preserving (sizeof(T)==sizeof(TOut))
151 }
152
153 std::unordered_map<std::string, size_t> getActualOutputSizesByName() const override {
154 return {{"output", actual_output_size_}};
155 }
156 size_t getActualOutputSize(int index) const override {
157 return (index == 0) ? actual_output_size_ : 0;
158 }
159
160 uint16_t getStageTypeId() const override {
161 return static_cast<uint16_t>(StageType::DIFFERENCE);
162 }
163
164 uint8_t getOutputDataType(size_t output_index) const override {
165 (void)output_index;
166 return static_cast<uint8_t>(getOutDataTypeEnum());
167 }
168
169 uint8_t getInputDataType(size_t /*input_index*/) const override {
170 return static_cast<uint8_t>(getInDataTypeEnum());
171 }
172
173 size_t serializeHeader(size_t output_index, uint8_t* buf, size_t max_size) const override {
174 (void)output_index;
175 // The mode byte only disambiguates fused (TOut != T) instantiations;
176 // same-type headers stay 6 bytes (preserves the legacy contract, see
177 // LegacyFloatHeaderCompatible).
178 constexpr size_t needed = std::is_same_v<T, TOut> ? 6 : 7;
179 if (max_size < needed) return 0;
180 buf[0] = static_cast<uint8_t>(getInDataTypeEnum());
181 buf[1] = static_cast<uint8_t>(getOutDataTypeEnum());
182 uint32_t cs = static_cast<uint32_t>(chunk_size_);
183 std::memcpy(buf + 2, &cs, sizeof(uint32_t));
184 if constexpr (needed == 7) buf[6] = static_cast<uint8_t>(Mode);
185 return needed;
186 }
187
188 void deserializeHeader(const uint8_t* buf, size_t size) override {
189 // DataTypes and FusionMode are baked into the template; factory picks
190 // the right instantiation. Only chunk_size needs to be restored at runtime.
191 if (size >= 6) {
192 uint32_t cs = 0;
193 std::memcpy(&cs, buf + 2, sizeof(uint32_t));
194 chunk_size_ = cs;
195 }
196 }
197
198 size_t getMaxHeaderSize(size_t output_index) const override {
199 (void)output_index;
200 return std::is_same_v<T, TOut> ? 6 : 7;
201 }
202
203 void saveState() override {
204 saved_chunk_size_ = chunk_size_;
205 saved_actual_output_size_ = actual_output_size_;
206 }
207
208 void restoreState() override {
209 chunk_size_ = saved_chunk_size_;
210 actual_output_size_ = saved_actual_output_size_;
211 }
212
213private:
214 size_t actual_output_size_;
215 size_t saved_actual_output_size_ = 0;
216 bool is_inverse_;
217 size_t chunk_size_;
218 size_t saved_chunk_size_ = 0;
219
220
221 DataType getInDataTypeEnum() const {
222 if (std::is_same_v<T, uint8_t>) return DataType::UINT8;
223 if (std::is_same_v<T, uint16_t>) return DataType::UINT16;
224 if (std::is_same_v<T, uint32_t>) return DataType::UINT32;
225 if (std::is_same_v<T, uint64_t>) return DataType::UINT64;
226 if (std::is_same_v<T, int8_t>) return DataType::INT8;
227 if (std::is_same_v<T, int16_t>) return DataType::INT16;
228 if (std::is_same_v<T, int32_t>) return DataType::INT32;
229 if (std::is_same_v<T, int64_t>) return DataType::INT64;
230 if (std::is_same_v<T, float>) return DataType::FLOAT32;
231 if (std::is_same_v<T, double>) return DataType::FLOAT64;
232 return DataType::UINT8;
233 }
234
235 DataType getOutDataTypeEnum() const {
236 if (std::is_same_v<TOut, uint8_t>) return DataType::UINT8;
237 if (std::is_same_v<TOut, uint16_t>) return DataType::UINT16;
238 if (std::is_same_v<TOut, uint32_t>) return DataType::UINT32;
239 if (std::is_same_v<TOut, uint64_t>) return DataType::UINT64;
240 if (std::is_same_v<TOut, int8_t>) return DataType::INT8;
241 if (std::is_same_v<TOut, int16_t>) return DataType::INT16;
242 if (std::is_same_v<TOut, int32_t>) return DataType::INT32;
243 if (std::is_same_v<TOut, int64_t>) return DataType::INT64;
244 if (std::is_same_v<TOut, float>) return DataType::FLOAT32;
245 if (std::is_same_v<TOut, double>) return DataType::FLOAT64;
246 return DataType::UINT8;
247 }
248};
249
250// ─── Same-type instantiations (original API, TOut = T) ───────────────────────
251extern template class DifferenceStage<float>;
252extern template class DifferenceStage<double>;
253extern template class DifferenceStage<int32_t>;
254extern template class DifferenceStage<int64_t>;
255extern template class DifferenceStage<uint16_t>;
256extern template class DifferenceStage<uint8_t>;
257extern template class DifferenceStage<uint32_t>;
258
259// ─── Negabinary-fused instantiations (TOut = unsigned counterpart of T; Mode defaults to NEGABINARY) ───
260extern template class DifferenceStage<int8_t, uint8_t>;
261extern template class DifferenceStage<int16_t, uint16_t>;
262extern template class DifferenceStage<int32_t, uint32_t>;
263extern template class DifferenceStage<int64_t, uint64_t>;
264
265// ─── Zigzag-fused instantiations (TOut = unsigned counterpart of T, Mode = ZIGZAG) ──────
266extern template class DifferenceStage<int8_t, uint8_t, FusionMode::ZIGZAG>;
267extern template class DifferenceStage<int16_t, uint16_t, FusionMode::ZIGZAG>;
268extern template class DifferenceStage<int32_t, uint32_t, FusionMode::ZIGZAG>;
269extern template class DifferenceStage<int64_t, uint64_t, FusionMode::ZIGZAG>;
270
271} // namespace fz
Definition diff.h:71
uint8_t getOutputDataType(size_t output_index) const override
Definition diff.h:164
uint8_t getInputDataType(size_t) const override
Definition diff.h:169
size_t getRequiredInputAlignment() const override
Definition diff.h:94
void setChunkSize(size_t bytes)
Definition diff.h:92
std::vector< size_t > estimateOutputSizes(const std::vector< size_t > &input_sizes) const override
Definition diff.h:147
size_t getActualOutputSize(int index) const override
Definition diff.h:156
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 saveState() override
Definition diff.h:203
void setInverse(bool inverse) override
Definition diff.h:82
FusedOpDecl getFusedOp() const override
Definition diff.h:127
uint16_t getStageTypeId() const override
Definition diff.h:160
size_t serializeHeader(size_t output_index, uint8_t *buf, size_t max_size) const override
Definition diff.h:173
size_t getMaxHeaderSize(size_t output_index) const override
Definition diff.h:198
void deserializeHeader(const uint8_t *buf, size_t size) override
Definition diff.h:188
std::unordered_map< std::string, size_t > getActualOutputSizesByName() const override
Definition diff.h:153
FusionSpec getFusionSpec() const override
Definition diff.h:114
std::string getName() const override
Definition diff.h:143
Definition mempool.h:82
Definition stage.h:31
FZM binary file format definitions — structs, enums, and helpers.
Definition dag.h:24
FusionMode
Definition diff.h:28
@ NEGABINARY
LC's DIFFNB — Negabinary<T>::encode/decode.
DataType
Element data type identifiers used in buffer and stage descriptors.
Definition fzm_format.h:142
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.