23#include <unordered_map>
44static_assert(
sizeof(LogTransformConfig) == 20,
45 "LogTransformConfig must be 20 bytes");
151template<
typename TInput =
float>
153 static_assert(std::is_floating_point<TInput>::value,
154 "LogTransformStage: TInput must be a floating-point type.");
190 explicit Config(
float eb,
float thr = 0.0f,
float cap = 0.05f)
196 ~LogTransformStage()
override;
200 const std::vector<void*>& inputs,
201 const std::vector<void*>& outputs,
202 const std::vector<size_t>& sizes)
override;
208 return sizeof(uint32_t);
215 std::string
getName()
const override {
return "LogTransform"; }
217 size_t getNumInputs()
const override {
return is_inverse_ ? 4 : 1; }
218 size_t getNumOutputs()
const override {
return is_inverse_ ? 1 : 4; }
221 if (is_inverse_)
return {
"reconstructed"};
222 return {
"output",
"signs",
"outlier_vals",
"outlier_idxs"};
226 const std::vector<size_t>& input_sizes)
const override;
228 std::unordered_map<std::string, size_t>
231 std::unordered_map<std::string, size_t> result;
232 for (
size_t i = 0; i < names.size() && i < actual_output_sizes_.size(); ++i)
233 result[names[i]] = actual_output_sizes_[i];
237 return (index >= 0 && index <
static_cast<int>(actual_output_sizes_.size()))
238 ? actual_output_sizes_[index] : 0;
241 void setInverse(
bool inverse)
override { is_inverse_ = inverse; }
242 bool isInverse()
const override {
return is_inverse_; }
249 if (is_inverse_)
return static_cast<uint8_t
>(inputDataType());
250 switch (output_index) {
251 case 0:
return static_cast<uint8_t
>(inputDataType());
252 case 1:
return static_cast<uint8_t
>(DataType::UINT8);
253 case 2:
return static_cast<uint8_t
>(inputDataType());
254 case 3:
return static_cast<uint8_t
>(DataType::UINT32);
255 default:
return static_cast<uint8_t
>(DataType::UINT8);
263 if (!is_inverse_)
return static_cast<uint8_t
>(inputDataType());
264 switch (input_index) {
265 case 0:
return static_cast<uint8_t
>(inputDataType());
266 case 1:
return static_cast<uint8_t
>(DataType::UINT8);
267 case 2:
return static_cast<uint8_t
>(inputDataType());
268 case 3:
return static_cast<uint8_t
>(DataType::UINT32);
269 default:
return static_cast<uint8_t
>(DataType::UINT8);
274 size_t serializeHeader(
size_t output_index, uint8_t* buf,
size_t max_size)
const override;
279 saved_config_ = config_;
280 saved_num_elements_ = num_elements_;
281 saved_actual_outlier_count_ = actual_outlier_count_;
282 saved_actual_output_sizes_ = actual_output_sizes_;
284 void restoreState()
override {
285 config_ = saved_config_;
286 num_elements_ = saved_num_elements_;
287 actual_outlier_count_ = saved_actual_outlier_count_;
288 actual_output_sizes_ = saved_actual_output_sizes_;
298 float getErrorBound()
const {
return config_.
error_bound; }
299 float getThreshold()
const {
return config_.
threshold; }
301 uint32_t getOutlierCount()
const {
return actual_outlier_count_; }
334 Config saved_config_;
335 bool is_inverse_ =
false;
337 std::vector<size_t> actual_output_sizes_{0, 0, 0, 0};
338 std::vector<size_t> saved_actual_output_sizes_{0, 0, 0, 0};
340 size_t num_elements_ = 0;
341 size_t saved_num_elements_ = 0;
342 uint32_t actual_outlier_count_ = 0;
343 uint32_t saved_actual_outlier_count_ = 0;
347 uint32_t* d_outlier_count_scratch_ =
nullptr;
348 MemoryPool* persistent_pool_ =
nullptr;
353 std::weak_ptr<const void> persistent_pool_alive_;
356 void initOutlierCountScratch(MemoryPool* pool);
369 size_t maxOutlierCount(
size_t n)
const {
370 const size_t scaled =
371 static_cast<size_t>(
static_cast<double>(n) * config_.outlier_capacity);
373 return (floored > n) ? n : floored;
376 static size_t signBytes(
size_t n) {
return (n + 7) / 8; }
379 return std::is_same<TInput, double>::value ? DataType::FLOAT64
384extern template class LogTransformStage<float>;
@ 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:139
Base class interface for all compression stages.
Backend-neutral GPU type aliases.