70template<
typename T =
float,
typename TOut = T, FusionMode Mode = FusionMode::NEGABINARY>
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).");
80 DifferenceStage() : actual_output_size_(0), is_inverse_(
false), chunk_size_(0) {}
82 void setInverse(
bool inverse)
override { is_inverse_ = inverse; }
83 bool isInverse()
const override {
return is_inverse_; }
93 size_t getChunkSize()
const {
return chunk_size_; }
95 return chunk_size_ > 0 ? chunk_size_ : 1;
102 if (is_inverse_ || chunk_size_ == 0)
return {};
104 return FusionSpec{FusionAccess::BlockLocal,
static_cast<uint32_t
>(chunk_size_)};
113 return FusedOpDecl{FusionStrategy::ChunkCooperative,
"DiffNegabinary",
114 "fused/chunk_fusion/chunk_fusion.cuh", {}};
120 const std::vector<void*>& inputs,
121 const std::vector<void*>& outputs,
122 const std::vector<size_t>& sizes
126 std::string
getName()
const override {
return "Difference"; }
127 size_t getNumInputs()
const override {
return 1; }
128 size_t getNumOutputs()
const override {
return 1; }
131 const std::vector<size_t>& input_sizes
133 return {input_sizes[0]};
137 return {{
"output", actual_output_size_}};
140 return (index == 0) ? actual_output_size_ : 0;
144 return static_cast<uint16_t
>(StageType::DIFFERENCE);
149 return static_cast<uint8_t
>(getOutDataTypeEnum());
153 return static_cast<uint8_t
>(getInDataTypeEnum());
156 size_t serializeHeader(
size_t output_index, uint8_t* buf,
size_t max_size)
const override {
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);
176 std::memcpy(&cs, buf + 2,
sizeof(uint32_t));
183 return std::is_same_v<T, TOut> ? 6 : 7;
187 saved_chunk_size_ = chunk_size_;
188 saved_actual_output_size_ = actual_output_size_;
191 void restoreState()
override {
192 chunk_size_ = saved_chunk_size_;
193 actual_output_size_ = saved_actual_output_size_;
197 size_t actual_output_size_;
198 size_t saved_actual_output_size_ = 0;
201 size_t saved_chunk_size_ = 0;
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;
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;
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>;
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>;
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>;
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
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.