diff --git a/kll/include/kll_helper_impl.hpp b/kll/include/kll_helper_impl.hpp index 31534d9a..56b5be69 100644 --- a/kll/include/kll_helper_impl.hpp +++ b/kll/include/kll_helper_impl.hpp @@ -39,8 +39,9 @@ uint8_t kll_helper::floor_of_log2_of_fraction(uint64_t numer, uint64_t denom) { if (denom > numer) { return 0; } uint8_t count = 0; while (true) { + // denom * 2 > numer, without overflowing denom + if (denom > (numer >> 1)) { return count; } denom <<= 1; - if (denom > numer) { return count; } count++; } } diff --git a/kll/include/kll_sketch.hpp b/kll/include/kll_sketch.hpp index d672c419..0133b899 100644 --- a/kll/include/kll_sketch.hpp +++ b/kll/include/kll_sketch.hpp @@ -237,6 +237,18 @@ class kll_sketch { template void update(FwdT&& item); + /** + * Updates this sketch with the given data item repeated the given number of times. + * The result is equivalent to calling update(item) weight times, at a cost that grows + * with the logarithm of the weight rather than with the weight itself. + * If cross-language portability is required, callers should ensure that + * the input string uses a compatible encoding (valid UTF-8). + * @param item from a stream of items + * @param weight number of times the item is repeated, must be positive + */ + template + void update(FwdT&& item, uint64_t weight); + /** * Merges another sketch into this one. * If sketches contain strings, callers are responsible for ensuring that @@ -570,6 +582,9 @@ class kll_sketch { std::unique_ptr items, uint32_t items_size, optional&& min_item, optional&& max_item, bool is_level_zero_sorted, const C& comparator); + // for weighted update + kll_sketch(uint16_t k, const T& item, uint64_t weight, const C& comparator, const A& allocator); + // common update code inline void update_min_max(const T& item); inline uint32_t internal_update(); diff --git a/kll/include/kll_sketch_impl.hpp b/kll/include/kll_sketch_impl.hpp index b12a39c8..fcc51d3a 100644 --- a/kll/include/kll_sketch_impl.hpp +++ b/kll/include/kll_sketch_impl.hpp @@ -186,6 +186,51 @@ void kll_sketch::update(FwdT&& item) { reset_sorted_view(); } +template +template +void kll_sketch::update(FwdT&& item, uint64_t weight) { + if (weight == 0) throw std::invalid_argument("weight must be positive"); + if (!check_update_item(item)) { return; } + if (weight < levels_[0]) { // fits into level zero without compaction + update_min_max(static_cast(item)); + for (uint64_t i = 1; i < weight; ++i) new (&items_[internal_update()]) T(static_cast(item)); + new (&items_[internal_update()]) T(std::forward(item)); + reset_sorted_view(); + } else { + merge(kll_sketch(k_, static_cast(item), weight, comparator_, allocator_)); + } +} + +// An item at level h carries weight 2^h, so one copy of the item at each level +// whose bit is set in the weight is an exact sketch of that many copies. +// Level capacities are defined up to 61 levels, so bits above the top level +// fold into weight >> 60 copies at level 60. +template +kll_sketch::kll_sketch(uint16_t k, const T& item, uint64_t weight, const C& comparator, const A& allocator): +comparator_(comparator), +allocator_(allocator), +k_(k), +m_(kll_constants::DEFAULT_M), +min_k_(k), +num_levels_(std::min(64 - count_leading_zeros_in_u64(weight), 61)), +is_level_zero_sorted_(true), +n_(weight), +levels_(num_levels_ + 1, 0, allocator), +items_(nullptr), +items_size_(0), +min_item_(item), +max_item_(item), +sorted_view_(nullptr) +{ + for (uint8_t level = 0; level < num_levels_; ++level) { + const uint64_t count = level + 1 < num_levels_ ? (weight >> level) & 1 : weight >> level; + levels_[level + 1] = levels_[level] + static_cast(count); + } + items_size_ = levels_[num_levels_]; + items_ = allocator_.allocate(items_size_); + for (uint32_t i = 0; i < items_size_; ++i) new (&items_[i]) T(item); +} + template void kll_sketch::update_min_max(const T& item) { if (is_empty()) { diff --git a/kll/test/kll_sketch_test.cpp b/kll/test/kll_sketch_test.cpp index b8d0b031..d4aa396e 100644 --- a/kll/test/kll_sketch_test.cpp +++ b/kll/test/kll_sketch_test.cpp @@ -443,6 +443,9 @@ TEST_CASE("kll sketch", "[kll_sketch]") { REQUIRE(kll_helper::floor_of_log2_of_fraction(6, 2) == 1); REQUIRE(kll_helper::floor_of_log2_of_fraction(7, 2) == 1); REQUIRE(kll_helper::floor_of_log2_of_fraction(8, 2) == 2); + REQUIRE(kll_helper::floor_of_log2_of_fraction(1ULL << 63, 1) == 63); + REQUIRE(kll_helper::floor_of_log2_of_fraction(std::numeric_limits::max(), 1) == 63); + REQUIRE(kll_helper::floor_of_log2_of_fraction(std::numeric_limits::max(), 3) == 62); } SECTION("out of order split points, float") { @@ -831,6 +834,123 @@ TEST_CASE("kll sketch", "[kll_sketch]") { REQUIRE(sb.get_n() == 3); } + SECTION("weighted update: zero weight") { + kll_float_sketch sketch(200, std::less(), 0); + REQUIRE_THROWS_AS(sketch.update(1.0f, 0), std::invalid_argument); + REQUIRE(sketch.is_empty()); + } + + SECTION("weighted update: NaN") { + kll_float_sketch sketch(200, std::less(), 0); + sketch.update(std::numeric_limits::quiet_NaN(), 1000); + REQUIRE(sketch.is_empty()); + } + + SECTION("weighted update: same as repeated updates in exact mode") { + kll_float_sketch weighted(200, std::less(), 0); + kll_float_sketch repeated(200, std::less(), 0); + for (int i = 1; i <= 15; ++i) { + weighted.update(static_cast(i), i); + for (int j = 0; j < i; ++j) repeated.update(static_cast(i)); + } + REQUIRE(weighted.get_n() == repeated.get_n()); + REQUIRE(weighted.get_num_retained() == repeated.get_num_retained()); + REQUIRE_FALSE(weighted.is_estimation_mode()); + for (int i = 0; i <= 16; ++i) { + REQUIRE(weighted.get_rank(static_cast(i)) == repeated.get_rank(static_cast(i))); + REQUIRE(weighted.get_rank(static_cast(i), false) == repeated.get_rank(static_cast(i), false)); + } + } + + SECTION("weighted update: one item with large weight") { + kll_float_sketch sketch(200, std::less(), 0); + const uint64_t weight = (1ULL << 40) + 12345; + sketch.update(7.0f, weight); + REQUIRE(sketch.get_n() == weight); + REQUIRE(sketch.get_num_retained() == 7); // number of bits set in the weight + REQUIRE(sketch.get_min_item() == 7.0f); + REQUIRE(sketch.get_max_item() == 7.0f); + REQUIRE(sketch.get_rank(7.0f) == 1); + REQUIRE(sketch.get_rank(7.0f, false) == 0); + REQUIRE(sketch.get_quantile(0.5) == 7.0f); + + uint64_t total_weight = 0; + for (auto pair: sketch) total_weight += pair.second; + REQUIRE(total_weight == weight); + } + + SECTION("weighted update: maximum weight") { + kll_float_sketch sketch(200, std::less(), 0); + sketch.update(1.0f); + sketch.update(2.0f, std::numeric_limits::max() - 1); + REQUIRE(sketch.get_n() == std::numeric_limits::max()); + REQUIRE(sketch.get_min_item() == 1.0f); + REQUIRE(sketch.get_max_item() == 2.0f); + REQUIRE(sketch.get_quantile(0.5) == 2.0f); + + uint64_t total_weight = 0; + for (auto pair: sketch) total_weight += pair.second; + REQUIRE(total_weight == sketch.get_n()); + + auto bytes = sketch.serialize(); + auto sketch2 = kll_float_sketch::deserialize(bytes.data(), bytes.size(), serde(), std::less(), 0); + REQUIRE(sketch2.get_n() == sketch.get_n()); + REQUIRE(sketch2.get_num_retained() == sketch.get_num_retained()); + } + + SECTION("weighted update: large weights into a full sketch") { + kll_float_sketch sketch(200, std::less(), 0); + for (int i = 0; i < 1000; ++i) sketch.update(static_cast(i)); + sketch.update(-1.0f, 1000); + sketch.update(2000.0f, 3000); + REQUIRE(sketch.get_n() == 5000); + REQUIRE(sketch.get_min_item() == -1.0f); + REQUIRE(sketch.get_max_item() == 2000.0f); + REQUIRE(sketch.get_rank(-1.0f) == Approx(0.2).margin(RANK_EPS_FOR_K_200)); + REQUIRE(sketch.get_rank(2000.0f, false) == Approx(0.4).margin(RANK_EPS_FOR_K_200)); + } + + SECTION("weighted update: estimation mode") { + kll_float_sketch sketch(200, std::less(), 0); + const int n = 10000; + std::vector weights(n); + uint64_t total = 0; + for (int i = 0; i < n; ++i) { + weights[i] = static_cast(i % 13) * 1000 + 1; + total += weights[i]; + sketch.update(static_cast(i), weights[i]); + } + REQUIRE(sketch.get_n() == total); + REQUIRE(sketch.is_estimation_mode()); + REQUIRE(sketch.get_min_item() == 0); + REQUIRE(sketch.get_max_item() == n - 1); + + uint64_t weight_below = 0; + for (int i = 0; i < n; ++i) { + const double true_rank = static_cast(weight_below) / total; + REQUIRE(sketch.get_rank(static_cast(i), false) == Approx(true_rank).margin(RANK_EPS_FOR_K_200)); + weight_below += weights[i]; + } + + auto bytes = sketch.serialize(); + auto sketch2 = kll_float_sketch::deserialize(bytes.data(), bytes.size(), serde(), std::less(), 0); + REQUIRE(sketch2.get_n() == sketch.get_n()); + REQUIRE(sketch2.get_num_retained() == sketch.get_num_retained()); + REQUIRE(sketch2.get_quantile(0.5) == sketch.get_quantile(0.5)); + } + + SECTION("weighted update: strings") { + kll_string_sketch sketch(200, std::less(), 0); + sketch.update(std::string("b"), 3); + const std::string a("a"); + sketch.update(a, 1000); + sketch.update(std::string("c"), 1ULL << 20); + REQUIRE(sketch.get_n() == 3 + 1000 + (1ULL << 20)); + REQUIRE(sketch.get_min_item() == "a"); + REQUIRE(sketch.get_max_item() == "c"); + REQUIRE(sketch.get_rank("a") == Approx(1000.0 / sketch.get_n()).margin(RANK_EPS_FOR_K_200)); + } + // cleanup REQUIRE(test_allocator_total_bytes == 0); }