From 57df19ce46ea8c87a5bf892323f55377cd10b24b Mon Sep 17 00:00:00 2001 From: jaideeppyne Date: Wed, 30 Sep 2026 08:00:01 +0530 Subject: [PATCH] fix(tdigest): correct rank and quantile tails The left tail of get_rank() was not divided by the total weight, and the right tail of get_quantile() was added onto the maximum. A last centroid of weight 2 makes that division 0/0, so return the maximum. Interpolation now weights each centroid by the distance to the other one. --- tdigest/include/tdigest_impl.hpp | 13 ++-- tdigest/test/tdigest_test.cpp | 102 +++++++++++++++++++++++++++++++ 2 files changed, 111 insertions(+), 4 deletions(-) diff --git a/tdigest/include/tdigest_impl.hpp b/tdigest/include/tdigest_impl.hpp index 2fb677da..958c8c40 100644 --- a/tdigest/include/tdigest_impl.hpp +++ b/tdigest/include/tdigest_impl.hpp @@ -133,7 +133,7 @@ double tdigest::get_rank(T value) const { if (value < first_mean) { if (first_mean - min_ > 0) { if (value == min_) return 0.5 / centroids_weight_; - return (1.0 + (value - min_) / (first_mean - min_) * (centroids_.front().get_weight() / 2.0 - 1.0)); // ? + return (1.0 + (value - min_) / (first_mean - min_) * (centroids_.front().get_weight() / 2.0 - 1.0)) / centroids_weight_; } return 0; // should never happen } @@ -143,7 +143,7 @@ double tdigest::get_rank(T value) const { if (value > last_mean) { if (max_ - last_mean > 0) { if (value == max_) return 1.0 - 0.5 / centroids_weight_; - return 1.0 - ((1.0 + (max_ - value) / (max_ - last_mean) * (centroids_.back().get_weight() / 2.0 - 1.0)) / centroids_weight_); // ? + return 1.0 - ((1.0 + (max_ - value) / (max_ - last_mean) * (centroids_.back().get_weight() / 2.0 - 1.0)) / centroids_weight_); } return 1; // should never happen } @@ -193,7 +193,10 @@ T tdigest::get_quantile(double rank) const { } const double last_weight = centroids_.back().get_weight(); if (last_weight > 1 && centroids_weight_ - weight <= last_weight / 2.0) { - return max_ + (centroids_weight_ - weight - 1.0) / (last_weight / 2.0 - 1.0) * (max_ - centroids_.back().get_mean()); + // A last centroid of weight 2 makes the denominator zero. The only rank that + // reaches this branch is the one that returns the stored maximum. + if (last_weight == 2) return max_; + return max_ - (centroids_weight_ - weight - 1.0) / (last_weight / 2.0 - 1.0) * (max_ - centroids_.back().get_mean()); } // interpolate between extremes @@ -214,7 +217,9 @@ T tdigest::get_quantile(double rank) const { } const double w1 = weight - weight_so_far - left_weight; const double w2 = weight_so_far + dw - weight - right_weight; - return weighted_average(centroids_[i].get_mean(), w1, centroids_[i + 1].get_mean(), w2); + // Weight each centroid by the distance to the other one, so the estimate + // moves toward the nearer centroid. + return weighted_average(centroids_[i].get_mean(), w2, centroids_[i + 1].get_mean(), w1); } weight_so_far += dw; } diff --git a/tdigest/test/tdigest_test.cpp b/tdigest/test/tdigest_test.cpp index d187b749..30b1fc1a 100644 --- a/tdigest/test/tdigest_test.cpp +++ b/tdigest/test/tdigest_test.cpp @@ -18,10 +18,14 @@ */ #include +#include +#include #include +#include #include #include #include +#include #include "tdigest.hpp" @@ -626,4 +630,102 @@ TEST_CASE("deserialize stream rejects zero centroid weight", "[tdigest]") { REQUIRE_THROWS_AS(tdigest_double::deserialize(is), std::invalid_argument); } +TEST_CASE("quantiles are monotonic and stay within min and max", "[tdigest]") { + tdigest_double td(100); + for (int i = 0; i < 10000; ++i) td.update(i); + double previous = td.get_min_value(); + for (int i = 0; i <= 1000; ++i) { + const double quantile = td.get_quantile(i / 1000.0); + REQUIRE(quantile >= previous); + REQUIRE(quantile >= td.get_min_value()); + REQUIRE(quantile <= td.get_max_value()); + previous = quantile; + } +} + +TEST_CASE("rank below the first centroid stays normalized", "[tdigest]") { + // The format allows a first centroid heavier than 1. The left tail of get_rank() + // must be divided by the total weight, the same way the right tail is. + tdigest_double source(200); + for (int i = 0; i < 1000; ++i) source.update(i); + auto bytes = source.serialize(); + double first_mean = 0; + std::memcpy(&first_mean, bytes.data() + first_centroid_mean_offset, sizeof(double)); + REQUIRE(first_mean == 0); + write_bytes(bytes, min_offset, -1.0); + write_bytes(bytes, first_centroid_weight_offset, static_cast(100)); + const auto td = tdigest_double::deserialize(bytes.data(), bytes.size()); + const double total_weight = static_cast(td.get_total_weight()); + REQUIRE(td.get_rank(-1) == 0.5 / total_weight); + REQUIRE(td.get_rank(-0.5) == (1.0 + ((100.0 / 2.0 - 1.0) * 0.5)) / total_weight); + double previous = 0; + for (int i = 0; i <= 100; ++i) { + const double rank = td.get_rank(-1.0 + (i / 100.0)); + REQUIRE(rank >= 0); + REQUIRE(rank <= 1); + REQUIRE(rank >= previous); + previous = rank; + } +} + +TEST_CASE("quantile above the last centroid does not exceed max", "[tdigest]") { + tdigest_double source(200); + for (int i = 0; i < 1000; ++i) source.update(i); + auto bytes = source.serialize(); + uint32_t num_centroids = 0; + std::memcpy(&num_centroids, bytes.data() + header_size, sizeof(uint32_t)); + REQUIRE(num_centroids > 1); + const size_t last_weight_offset = first_centroid_weight_offset + (num_centroids - 1) * 16; + write_bytes(bytes, last_weight_offset, static_cast(100)); + const auto td = tdigest_double::deserialize(bytes.data(), bytes.size()); + double previous = td.get_min_value(); + for (int i = 0; i <= 1000; ++i) { + const double quantile = td.get_quantile(i / 1000.0); + REQUIRE(quantile >= previous); + REQUIRE(quantile <= td.get_max_value()); + previous = quantile; + } +} + +tdigest_double sketch_from_centroids(double min, double max, + std::initializer_list> centroids) { + const auto num_centroids = static_cast(centroids.size()); + std::vector bytes(first_centroid_mean_offset + static_cast(num_centroids) * 16, 0); + bytes[0] = 2; // preamble longs + bytes[1] = 1; // serial version + bytes[2] = 20; // sketch type + write_bytes(bytes, 3, static_cast(100)); + write_bytes(bytes, header_size, num_centroids); + write_bytes(bytes, header_size + sizeof(uint32_t), static_cast(0)); + write_bytes(bytes, min_offset, min); + write_bytes(bytes, max_offset, max); + size_t offset = first_centroid_mean_offset; + for (const auto& centroid : centroids) { + write_bytes(bytes, offset, centroid.first); + write_bytes(bytes, offset + sizeof(double), centroid.second); + offset += 16; + } + return tdigest_double::deserialize(bytes.data(), bytes.size()); +} + +TEST_CASE("quantile with a last centroid of weight 2 returns max", "[tdigest]") { + // weight == total - 1 and last weight == 2 is 0/0 unless that case returns max. + const auto td = sketch_from_centroids(0, 20, {{0, 10}, {10, 2}}); + const double quantile = td.get_quantile(11.0 / 12.0); + REQUIRE_FALSE(std::isnan(quantile)); + REQUIRE(quantile == 20); +} + +TEST_CASE("quantile right tail approaches max from below", "[tdigest]") { + const auto td = sketch_from_centroids(0, 40, {{10, 100}, {20, 100}, {30, 100}}); + REQUIRE(td.get_quantile(0.9) == Approx(34.081632653061224).epsilon(1e-12)); +} + +TEST_CASE("quantile interpolation weights the nearer centroid more", "[tdigest]") { + const auto td = sketch_from_centroids(0, 40, {{10, 100}, {20, 100}, {30, 100}}); + // Target weight 80 sits between 10 and 20, closer to 10. + // weighted_average(10, 70, 20, 30) == 13; the swapped weights would return 17. + REQUIRE(td.get_quantile(80.0 / 300.0) == 13); +} + } /* namespace datasketches */