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 // Block-local (chunk-cooperative) when chunking a signed->unsigned negabinary
99 // difference — the fused DiffNegabinary op reproduces exactly this. block_size
100 // is the chunk in BYTES (the granularity the whole chunk chain shares).
101 FusionSpec getFusionSpec() const override {
102 if (is_inverse_ || chunk_size_ == 0) return {};
103 if constexpr (!std::is_same_v<T, TOut> && Mode == FusionMode::NEGABINARY)
104 return FusionSpec{FusionAccess::BlockLocal, static_cast<uint32_t>(chunk_size_)};
105 else
106 return {};
107 }
108
111 FusedOpDecl getFusedOp() const override {
112 if (!getFusionSpec().fusable()) return {};
113 return FusedOpDecl{FusionStrategy::ChunkCooperative, "DiffNegabinary",
114 "fused/chunk_fusion/chunk_fusion.cuh", {}};
115 }
116
118 fz::stream_t stream,
119 MemoryPool* pool,
120 const std::vector<void*>& inputs,
121 const std::vector<void*>& outputs,
122 const std::vector<size_t>& sizes
123 ) override;
124
125
126 std::string getName() const override { return "Difference"; }
127 size_t getNumInputs() const override { return 1; }
128 size_t getNumOutputs() const override { return 1; }
129
130 std::vector<size_t> estimateOutputSizes(
131 const std::vector<size_t>& input_sizes
132 ) const override {
133 return {input_sizes[0]}; // size-preserving (sizeof(T)==sizeof(TOut))
134 }
135
136 std::unordered_map<std::string, size_t> getActualOutputSizesByName() const override {
137 return {{"output", actual_output_size_}};
138 }
139 size_t getActualOutputSize(int index) const override {
140 return (index == 0) ? actual_output_size_ : 0;
141 }
142
143 uint16_t getStageTypeId() const override {
144 return static_cast<uint16_t>(StageType::DIFFERENCE);
145 }
146
147 uint8_t getOutputDataType(size_t output_index) const override {
148 (void)output_index;
149 return static_cast<uint8_t>(getOutDataTypeEnum());
150 }
151
152 uint8_t getInputDataType(size_t /*input_index*/) const override {
153 return static_cast<uint8_t>(getInDataTypeEnum());
154 }
155
156 size_t serializeHeader(size_t output_index, uint8_t* buf, size_t max_size) const override {
157 (void)output_index;
158 // The mode byte only disambiguates fused (TOut != T) instantiations;
159 // same-type headers stay 6 bytes (preserves the legacy contract, see
160 // LegacyFloatHeaderCompatible).
161 constexpr size_t needed = std::is_same_v<T, TOut> ? 6 : 7;
162 if (max_size < needed) return 0;
163 buf[0] = static_cast<uint8_t>(getInDataTypeEnum());
164 buf[1] = static_cast<uint8_t>(getOutDataTypeEnum());
165 uint32_t cs = static_cast<uint32_t>(chunk_size_);
166 std::memcpy(buf + 2, &cs, sizeof(uint32_t));
167 if constexpr (needed == 7) buf[6] = static_cast<uint8_t>(Mode);
168 return needed;
169 }
170
171 void deserializeHeader(const uint8_t* buf, size_t size) override {
172 // DataTypes and FusionMode are baked into the template; factory picks
173 // the right instantiation. Only chunk_size needs to be restored at runtime.
174 if (size >= 6) {
175 uint32_t cs = 0;
176 std::memcpy(&cs, buf + 2, sizeof(uint32_t));
177 chunk_size_ = cs;
178 }
179 }
180
181 size_t getMaxHeaderSize(size_t output_index) const override {
182 (void)output_index;
183 return std::is_same_v<T, TOut> ? 6 : 7;
184 }
185
186 void saveState() override {
187 saved_chunk_size_ = chunk_size_;
188 saved_actual_output_size_ = actual_output_size_;
189 }
190
191 void restoreState() override {
192 chunk_size_ = saved_chunk_size_;
193 actual_output_size_ = saved_actual_output_size_;
194 }
195
196private:
197 size_t actual_output_size_;
198 size_t saved_actual_output_size_ = 0;
199 bool is_inverse_;
200 size_t chunk_size_;
201 size_t saved_chunk_size_ = 0;
202
203
204 DataType getInDataTypeEnum() const {
205 if (std::is_same_v<T, uint8_t>) return DataType::UINT8;
206 if (std::is_same_v<T, uint16_t>) return DataType::UINT16;
207 if (std::is_same_v<T, uint32_t>) return DataType::UINT32;
208 if (std::is_same_v<T, uint64_t>) return DataType::UINT64;
209 if (std::is_same_v<T, int8_t>) return DataType::INT8;
210 if (std::is_same_v<T, int16_t>) return DataType::INT16;
211 if (std::is_same_v<T, int32_t>) return DataType::INT32;
212 if (std::is_same_v<T, int64_t>) return DataType::INT64;
213 if (std::is_same_v<T, float>) return DataType::FLOAT32;
214 if (std::is_same_v<T, double>) return DataType::FLOAT64;
215 return DataType::UINT8;
216 }
217
218 DataType getOutDataTypeEnum() const {
219 if (std::is_same_v<TOut, uint8_t>) return DataType::UINT8;
220 if (std::is_same_v<TOut, uint16_t>) return DataType::UINT16;
221 if (std::is_same_v<TOut, uint32_t>) return DataType::UINT32;
222 if (std::is_same_v<TOut, uint64_t>) return DataType::UINT64;
223 if (std::is_same_v<TOut, int8_t>) return DataType::INT8;
224 if (std::is_same_v<TOut, int16_t>) return DataType::INT16;
225 if (std::is_same_v<TOut, int32_t>) return DataType::INT32;
226 if (std::is_same_v<TOut, int64_t>) return DataType::INT64;
227 if (std::is_same_v<TOut, float>) return DataType::FLOAT32;
228 if (std::is_same_v<TOut, double>) return DataType::FLOAT64;
229 return DataType::UINT8;
230 }
231};
232
233// ─── Same-type instantiations (original API, TOut = T) ───────────────────────
234extern template class DifferenceStage<float>;
235extern template class DifferenceStage<double>;
236extern template class DifferenceStage<int32_t>;
237extern template class DifferenceStage<int64_t>;
238extern template class DifferenceStage<uint16_t>;
239extern template class DifferenceStage<uint8_t>;
240extern template class DifferenceStage<uint32_t>;
241
242// ─── Negabinary-fused instantiations (TOut = unsigned counterpart of T; Mode defaults to NEGABINARY) ───
243extern template class DifferenceStage<int8_t, uint8_t>;
244extern template class DifferenceStage<int16_t, uint16_t>;
245extern template class DifferenceStage<int32_t, uint32_t>;
246extern template class DifferenceStage<int64_t, uint64_t>;
247
248// ─── Zigzag-fused instantiations (TOut = unsigned counterpart of T, Mode = ZIGZAG) ──────
249extern template class DifferenceStage<int8_t, uint8_t, FusionMode::ZIGZAG>;
250extern template class DifferenceStage<int16_t, uint16_t, FusionMode::ZIGZAG>;
251extern template class DifferenceStage<int32_t, uint32_t, FusionMode::ZIGZAG>;
252extern template class DifferenceStage<int64_t, uint64_t, FusionMode::ZIGZAG>;
253
254} // namespace fz
Definition diff.h:71
uint8_t getOutputDataType(size_t output_index) const override
Definition diff.h:147
uint8_t getInputDataType(size_t) const override
Definition diff.h:152
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:130
size_t getActualOutputSize(int index) const override
Definition diff.h:139
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:186
void setInverse(bool inverse) override
Definition diff.h:82
FusedOpDecl getFusedOp() const override
Definition diff.h:111
uint16_t getStageTypeId() const override
Definition diff.h:143
size_t serializeHeader(size_t output_index, uint8_t *buf, size_t max_size) const override
Definition diff.h:156
size_t getMaxHeaderSize(size_t output_index) const override
Definition diff.h:181
void deserializeHeader(const uint8_t *buf, size_t size) override
Definition diff.h:171
std::unordered_map< std::string, size_t > getActualOutputSizesByName() const override
Definition diff.h:136
FusionSpec getFusionSpec() const override
Definition diff.h:101
std::string getName() const override
Definition diff.h:126
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:139
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:52
Backend-neutral GPU type aliases.