Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 9 additions & 4 deletions tdigest/include/tdigest_impl.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -133,7 +133,7 @@ double tdigest<T, A>::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
}
Expand All @@ -143,7 +143,7 @@ double tdigest<T, A>::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
}
Expand Down Expand Up @@ -193,7 +193,10 @@ T tdigest<T, A>::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
Expand All @@ -214,7 +217,9 @@ T tdigest<T, A>::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;
}
Expand Down
102 changes: 102 additions & 0 deletions tdigest/test/tdigest_test.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -18,10 +18,14 @@
*/

#include <catch2/catch.hpp>
#include <cmath>
#include <cstdint>
#include <cstring>
#include <initializer_list>
#include <iostream>
#include <fstream>
#include <sstream>
#include <utility>

#include "tdigest.hpp"

Expand Down Expand Up @@ -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<uint64_t>(100));
const auto td = tdigest_double::deserialize(bytes.data(), bytes.size());
const double total_weight = static_cast<double>(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<uint64_t>(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<std::pair<double, uint64_t>> centroids) {
const auto num_centroids = static_cast<uint32_t>(centroids.size());
std::vector<uint8_t> bytes(first_centroid_mean_offset + static_cast<size_t>(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<uint16_t>(100));
write_bytes(bytes, header_size, num_centroids);
write_bytes(bytes, header_size + sizeof(uint32_t), static_cast<uint32_t>(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 */
Loading