23#include <unordered_map>
44static_assert(
sizeof(LogTransformConfig) == 20,
45 "LogTransformConfig must be 20 bytes");
150template<
typename TInput =
float>
152 static_assert(std::is_floating_point<TInput>::value,
153 "LogTransformStage: TInput must be a floating-point type.");
189 explicit Config(
float eb,
float thr = 0.0f,
float cap = 0.05f)
195 ~LogTransformStage()
override;
199 const std::vector<void*>& inputs,
200 const std::vector<void*>& outputs,
201 const std::vector<size_t>& sizes)
override;
207 return sizeof(uint32_t);
214 std::string
getName()
const override {
return "LogTransform"; }
216 size_t getNumInputs()
const override {
return is_inverse_ ? 4 : 1; }
217 size_t getNumOutputs()
const override {
return is_inverse_ ? 1 : 4; }
220 if (is_inverse_)
return {
"reconstructed"};
221 return {
"output",
"signs",
"outlier_vals",
"outlier_idxs"};
225 const std::vector<size_t>& input_sizes)
const override;
227 std::unordered_map<std::string, size_t>
230 std::unordered_map<std::string, size_t> result;
231 for (
size_t i = 0; i < names.size() && i < actual_output_sizes_.size(); ++i)
232 result[names[i]] = actual_output_sizes_[i];
236 return (index >= 0 && index <
static_cast<int>(actual_output_sizes_.size()))
237 ? actual_output_sizes_[index] : 0;
240 void setInverse(
bool inverse)
override { is_inverse_ = inverse; }
241 bool isInverse()
const override {
return is_inverse_; }
248 if (is_inverse_)
return static_cast<uint8_t
>(inputDataType());
249 switch (output_index) {
250 case 0:
return static_cast<uint8_t
>(inputDataType());
251 case 1:
return static_cast<uint8_t
>(DataType::UINT8);
252 case 2:
return static_cast<uint8_t
>(inputDataType());
253 case 3:
return static_cast<uint8_t
>(DataType::UINT32);
254 default:
return static_cast<uint8_t
>(DataType::UINT8);
262 if (!is_inverse_)
return static_cast<uint8_t
>(inputDataType());
263 switch (input_index) {
264 case 0:
return static_cast<uint8_t
>(inputDataType());
265 case 1:
return static_cast<uint8_t
>(DataType::UINT8);
266 case 2:
return static_cast<uint8_t
>(inputDataType());
267 case 3:
return static_cast<uint8_t
>(DataType::UINT32);
268 default:
return static_cast<uint8_t
>(DataType::UINT8);
273 size_t serializeHeader(
size_t output_index, uint8_t* buf,
size_t max_size)
const override;
278 saved_config_ = config_;
279 saved_num_elements_ = num_elements_;
280 saved_actual_outlier_count_ = actual_outlier_count_;
281 saved_actual_output_sizes_ = actual_output_sizes_;
283 void restoreState()
override {
284 config_ = saved_config_;
285 num_elements_ = saved_num_elements_;
286 actual_outlier_count_ = saved_actual_outlier_count_;
287 actual_output_sizes_ = saved_actual_output_sizes_;
297 float getErrorBound()
const {
return config_.
error_bound; }
298 float getThreshold()
const {
return config_.
threshold; }
300 uint32_t getOutlierCount()
const {
return actual_outlier_count_; }
333 Config saved_config_;
334 bool is_inverse_ =
false;
336 std::vector<size_t> actual_output_sizes_{0, 0, 0, 0};
337 std::vector<size_t> saved_actual_output_sizes_{0, 0, 0, 0};
339 size_t num_elements_ = 0;
340 size_t saved_num_elements_ = 0;
341 uint32_t actual_outlier_count_ = 0;
342 uint32_t saved_actual_outlier_count_ = 0;
346 uint32_t* d_outlier_count_scratch_ =
nullptr;
347 MemoryPool* persistent_pool_ =
nullptr;
352 std::weak_ptr<const void> persistent_pool_alive_;
355 void initOutlierCountScratch(MemoryPool* pool);
368 size_t maxOutlierCount(
size_t n)
const {
369 const size_t scaled =
370 static_cast<size_t>(
static_cast<double>(n) * config_.outlier_capacity);
372 return (floored > n) ? n : floored;
375 static size_t signBytes(
size_t n) {
return (n + 7) / 8; }
378 return std::is_same<TInput, double>::value ? DataType::FLOAT64
383extern template class LogTransformStage<float>;
Definition algorithms.h:48
@ LOG_TRANSFORM
Log-space transform for point-wise relative bounds (Liang et al., CLUSTER'18)
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.