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 void execute(
99 fz::stream_t stream,
100 MemoryPool* pool,
101 const std::vector<void*>& inputs,
102 const std::vector<void*>& outputs,
103 const std::vector<size_t>& sizes
104 ) override;
105
106
107 std::string getName() const override { return "Difference"; }
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 return {input_sizes[0]}; // size-preserving (sizeof(T)==sizeof(TOut))
115 }
116
117 std::unordered_map<std::string, size_t> getActualOutputSizesByName() const override {
118 return {{"output", actual_output_size_}};
119 }
120 size_t getActualOutputSize(int index) const override {
121 return (index == 0) ? actual_output_size_ : 0;
122 }
123
124 uint16_t getStageTypeId() const override {
125 return static_cast<uint16_t>(StageType::DIFFERENCE);
126 }
127
128 uint8_t getOutputDataType(size_t output_index) const override {
129 (void)output_index;
130 return static_cast<uint8_t>(getOutDataTypeEnum());
131 }
132
133 uint8_t getInputDataType(size_t /*input_index*/) const override {
134 return static_cast<uint8_t>(getInDataTypeEnum());
135 }
136
137 size_t serializeHeader(size_t output_index, uint8_t* buf, size_t max_size) const override {
138 (void)output_index;
139 // The mode byte only disambiguates fused (TOut != T) instantiations;
140 // same-type headers stay 6 bytes (preserves the legacy contract, see
141 // LegacyFloatHeaderCompatible).
142 constexpr size_t needed = std::is_same_v<T, TOut> ? 6 : 7;
143 if (max_size < needed) return 0;
144 buf[0] = static_cast<uint8_t>(getInDataTypeEnum());
145 buf[1] = static_cast<uint8_t>(getOutDataTypeEnum());
146 uint32_t cs = static_cast<uint32_t>(chunk_size_);
147 std::memcpy(buf + 2, &cs, sizeof(uint32_t));
148 if constexpr (needed == 7) buf[6] = static_cast<uint8_t>(Mode);
149 return needed;
150 }
151
152 void deserializeHeader(const uint8_t* buf, size_t size) override {
153 // DataTypes and FusionMode are baked into the template; factory picks
154 // the right instantiation. Only chunk_size needs to be restored at runtime.
155 if (size >= 6) {
156 uint32_t cs = 0;
157 std::memcpy(&cs, buf + 2, sizeof(uint32_t));
158 chunk_size_ = cs;
159 }
160 }
161
162 size_t getMaxHeaderSize(size_t output_index) const override {
163 (void)output_index;
164 return std::is_same_v<T, TOut> ? 6 : 7;
165 }
166
167 void saveState() override {
168 saved_chunk_size_ = chunk_size_;
169 saved_actual_output_size_ = actual_output_size_;
170 }
171
172 void restoreState() override {
173 chunk_size_ = saved_chunk_size_;
174 actual_output_size_ = saved_actual_output_size_;
175 }
176
177private:
178 size_t actual_output_size_;
179 size_t saved_actual_output_size_ = 0;
180 bool is_inverse_;
181 size_t chunk_size_;
182 size_t saved_chunk_size_ = 0;
183
184
185 DataType getInDataTypeEnum() const {
186 if (std::is_same_v<T, uint8_t>) return DataType::UINT8;
187 if (std::is_same_v<T, uint16_t>) return DataType::UINT16;
188 if (std::is_same_v<T, uint32_t>) return DataType::UINT32;
189 if (std::is_same_v<T, uint64_t>) return DataType::UINT64;
190 if (std::is_same_v<T, int8_t>) return DataType::INT8;
191 if (std::is_same_v<T, int16_t>) return DataType::INT16;
192 if (std::is_same_v<T, int32_t>) return DataType::INT32;
193 if (std::is_same_v<T, int64_t>) return DataType::INT64;
194 if (std::is_same_v<T, float>) return DataType::FLOAT32;
195 if (std::is_same_v<T, double>) return DataType::FLOAT64;
196 return DataType::UINT8;
197 }
198
199 DataType getOutDataTypeEnum() const {
200 if (std::is_same_v<TOut, uint8_t>) return DataType::UINT8;
201 if (std::is_same_v<TOut, uint16_t>) return DataType::UINT16;
202 if (std::is_same_v<TOut, uint32_t>) return DataType::UINT32;
203 if (std::is_same_v<TOut, uint64_t>) return DataType::UINT64;
204 if (std::is_same_v<TOut, int8_t>) return DataType::INT8;
205 if (std::is_same_v<TOut, int16_t>) return DataType::INT16;
206 if (std::is_same_v<TOut, int32_t>) return DataType::INT32;
207 if (std::is_same_v<TOut, int64_t>) return DataType::INT64;
208 if (std::is_same_v<TOut, float>) return DataType::FLOAT32;
209 if (std::is_same_v<TOut, double>) return DataType::FLOAT64;
210 return DataType::UINT8;
211 }
212};
213
214// ─── Same-type instantiations (original API, TOut = T) ───────────────────────
215extern template class DifferenceStage<float>;
216extern template class DifferenceStage<double>;
217extern template class DifferenceStage<int32_t>;
218extern template class DifferenceStage<int64_t>;
219extern template class DifferenceStage<uint16_t>;
220extern template class DifferenceStage<uint8_t>;
221extern template class DifferenceStage<uint32_t>;
222
223// ─── Negabinary-fused instantiations (TOut = unsigned counterpart of T; Mode defaults to NEGABINARY) ───
224extern template class DifferenceStage<int8_t, uint8_t>;
225extern template class DifferenceStage<int16_t, uint16_t>;
226extern template class DifferenceStage<int32_t, uint32_t>;
227extern template class DifferenceStage<int64_t, uint64_t>;
228
229// ─── Zigzag-fused instantiations (TOut = unsigned counterpart of T, Mode = ZIGZAG) ──────
230extern template class DifferenceStage<int8_t, uint8_t, FusionMode::ZIGZAG>;
231extern template class DifferenceStage<int16_t, uint16_t, FusionMode::ZIGZAG>;
232extern template class DifferenceStage<int32_t, uint32_t, FusionMode::ZIGZAG>;
233extern template class DifferenceStage<int64_t, uint64_t, FusionMode::ZIGZAG>;
234
235} // namespace fz
Definition diff.h:71
uint8_t getOutputDataType(size_t output_index) const override
Definition diff.h:128
uint8_t getInputDataType(size_t) const override
Definition diff.h:133
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:111
size_t getActualOutputSize(int index) const override
Definition diff.h:120
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:167
void setInverse(bool inverse) override
Definition diff.h:82
uint16_t getStageTypeId() const override
Definition diff.h:124
size_t serializeHeader(size_t output_index, uint8_t *buf, size_t max_size) const override
Definition diff.h:137
size_t getMaxHeaderSize(size_t output_index) const override
Definition diff.h:162
void deserializeHeader(const uint8_t *buf, size_t size) override
Definition diff.h:152
std::unordered_map< std::string, size_t > getActualOutputSizesByName() const override
Definition diff.h:117
std::string getName() const override
Definition diff.h:107
Definition mempool.h:82
Definition stage.h:30
FZM binary file format definitions — structs, enums, and helpers.
Definition algorithms.h:48
FusionMode
Definition diff.h:28
DataType
Element data type identifiers used in buffer and stage descriptors.
Definition fzm_format.h:117
Base class interface for all compression stages.
Backend-neutral GPU type aliases.