Skip to content
Open
Show file tree
Hide file tree
Changes from 4 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
42 changes: 41 additions & 1 deletion src/commands/cmd_tdigest.cc
Original file line number Diff line number Diff line change
Expand Up @@ -556,6 +556,45 @@ class CommandTDigestTrimmedMean : public Commander {
double high_cut_quantile_;
};

class CommandTDigestCDF : public Commander {
Status Parse(const std::vector<std::string> &args) override {
if (args.size() == 2) return {Status::RedisParseErr, errWrongNumOfArguments};
key_name_ = args[1];
inputs_.reserve(args.size() - 2);
for (size_t i = 2; i < args.size(); i++) {
auto value = ParseFloat(args[i]);
if (!value) {
return {Status::RedisParseErr, errValueIsNotFloat};
}
if (std::isnan(*value)) {
return {Status::RedisParseErr, errValueIsNotFloat};
}
inputs_.push_back(*value);
}
return Status::OK();
}

Status Execute(engine::Context &ctx, Server *srv, Connection *conn, std::string *output) override {
TDigest tdigest(srv->storage, conn->GetNamespace());
TDigestCDFResult result;
auto s = tdigest.CDF(ctx, key_name_, inputs_, &result);
if (!s.ok()) {
if (s.IsNotFound()) {
return {Status::RedisExecErr, errKeyNotFound};
}
return {Status::RedisExecErr, s.ToString()};
}

*output =
conn->MultiBulkString(result.cdf_values | ranges::views::transform(util::Float2String) | ranges::to_vector);
Comment thread
jihuayu marked this conversation as resolved.
Outdated
return Status::OK();
}

private:
std::string key_name_;
std::vector<double> inputs_;
};

std::vector<CommandKeyRange> GetMergeKeyRange(const std::vector<std::string> &args) {
auto numkeys = ParseInt<int>(args[2], 10).ValueOr(0);
return {{1, 1, 1}, {3, 2 + numkeys, 1}};
Expand All @@ -573,5 +612,6 @@ REDIS_REGISTER_COMMANDS(TDigest, MakeCmdAttr<CommandTDigestCreate>("tdigest.crea
MakeCmdAttr<CommandTDigestQuantile>("tdigest.quantile", -3, "read-only", 1, 1, 1),
MakeCmdAttr<CommandTDigestTrimmedMean>("tdigest.trimmed_mean", 4, "read-only", 1, 1, 1),
MakeCmdAttr<CommandTDigestReset>("tdigest.reset", 2, "write", 1, 1, 1),
MakeCmdAttr<CommandTDigestMerge>("tdigest.merge", -4, "write", GetMergeKeyRange));
MakeCmdAttr<CommandTDigestMerge>("tdigest.merge", -4, "write", GetMergeKeyRange),
MakeCmdAttr<CommandTDigestCDF>("tdigest.cdf", -3, "read-only", 1, 1, 1));
} // namespace redis
89 changes: 89 additions & 0 deletions src/types/redis_tdigest.cc
Original file line number Diff line number Diff line change
Expand Up @@ -28,13 +28,16 @@
#include <rocksdb/status.h>

#include <algorithm>
#include <cstdint>
#include <iterator>
#include <limits>
#include <memory>
#include <range/v3/algorithm/minmax.hpp>
#include <range/v3/range/conversion.hpp>
#include <range/v3/view/join.hpp>
#include <range/v3/view/map.hpp>
#include <range/v3/view/transform.hpp>
#include <set>
#include <vector>

#include "commands/error_constants.h"
Expand Down Expand Up @@ -570,6 +573,92 @@ rocksdb::Status TDigest::Merge(engine::Context& ctx, const Slice& dest_digest,
return storage_->Write(ctx, storage_->DefaultWriteOptions(), batch->GetWriteBatch());
}

rocksdb::Status TDigest::CDF(engine::Context& ctx, const Slice& digest_name, const std::vector<double>& inputs,
TDigestCDFResult* result) {
std::map<double, std::vector<size_t>> sorted_unique_inputs_with_idx;
for (size_t i = 0; i < inputs.size(); ++i) {
sorted_unique_inputs_with_idx[inputs[i]].push_back(i);
}

std::vector<double> cdf_values;
if (auto status = cdfUniqSorted(ctx, digest_name,
sorted_unique_inputs_with_idx | ranges::views::keys | ranges::to_vector, &cdf_values);
!status.ok()) {
return status;
}
result->cdf_values.resize(inputs.size(), std::numeric_limits<double>::quiet_NaN());

size_t idx = 0;
for (auto iter = sorted_unique_inputs_with_idx.cbegin(); iter != sorted_unique_inputs_with_idx.cend(); ++iter) {
for (auto original_idx : iter->second) {
result->cdf_values[original_idx] = cdf_values[idx];
}
++idx;
}

return rocksdb::Status::OK();
}

rocksdb::Status TDigest::cdfUniqSorted(engine::Context& ctx, const Slice& digest_name,
const std::vector<double>& inputs, std::vector<double>* values) {
auto ns_key = AppendNamespacePrefix(digest_name);
TDigestMetadata metadata;
{
LockGuard guard(storage_->GetLockManager(), ns_key);

if (auto status = getMetaDataByNsKey(ctx, ns_key, &metadata); !status.ok()) {
return status;
}

if (metadata.total_observations == 0) {
*values = std::vector<double>(inputs.size(), std::numeric_limits<double>::quiet_NaN());
return rocksdb::Status::OK();
}

if (auto status = mergeNodes(ctx, ns_key, &metadata); !status.ok()) {
return status;
}
}

std::vector<Centroid> centroids;
if (auto status = dumpCentroids(ctx, ns_key, metadata, &centroids); !status.ok()) {
return status;
}

auto dump_centroids = DummyCentroids<false>(metadata, centroids);
auto total_weight = dump_centroids.TotalWeight();
auto iter = dump_centroids.Begin();
double accum_weight = 0.;
std::vector<double> results;
results.reserve(inputs.size());
for (const auto val : inputs) {
double weight = accum_weight;
for (; iter->Valid(); iter->Next()) {
auto current_centroid_result = iter->GetCentroid();
if (!current_centroid_result) {
return rocksdb::Status::InvalidArgument(current_centroid_result.Msg());
}
auto& current_centroid = *current_centroid_result;

const int cmp = DoubleCompare(val, current_centroid.mean);

if (cmp < 0) {
break;
}
accum_weight += current_centroid.weight;
if (cmp > 0) {
weight += current_centroid.weight;
continue;
}
weight += current_centroid.weight / 2;
}
double cdf_val = (weight / total_weight);
results.push_back(cdf_val);
}
*values = std::move(results);
return rocksdb::Status::OK();
}

rocksdb::Status TDigest::GetMetaData(engine::Context& context, const Slice& digest_name, TDigestMetadata* metadata) {
auto ns_key = AppendNamespacePrefix(digest_name);
return Database::GetMetadata(context, {kRedisTDigest}, ns_key, metadata);
Expand Down
10 changes: 10 additions & 0 deletions src/types/redis_tdigest.h
Original file line number Diff line number Diff line change
Expand Up @@ -49,6 +49,10 @@ struct TDigestMergeOptions {
bool override_flag = false;
};

struct TDigestCDFResult {
std::vector<double> cdf_values;
};

struct TDigestQuantitleResult {
std::optional<std::vector<double>> quantiles;
};
Expand Down Expand Up @@ -93,6 +97,9 @@ class TDigest : public SubKeyScanner {
double high_cut_quantile, TDigestTrimmedMeanResult* result);
rocksdb::Status GetMetaData(engine::Context& context, const Slice& digest_name, TDigestMetadata* metadata);

rocksdb::Status CDF(engine::Context& ctx, const Slice& digest_name, const std::vector<double>& inputs,
TDigestCDFResult* result);

private:
enum class SegmentType : uint8_t { kBuffer = 0, kCentroids = 1, kGuardFlag = 0xFF };

Expand Down Expand Up @@ -160,5 +167,8 @@ class TDigest : public SubKeyScanner {
Centroid* centroid) const;
rocksdb::Status prepareRankData(engine::Context& ctx, const Slice& digest_name, TDigestMetadata& metadata,
std::vector<Centroid>& centroids);

rocksdb::Status cdfUniqSorted(engine::Context& ctx, const Slice& digest_name, const std::vector<double>& inputs,
std::vector<double>* values);
};
} // namespace redis
178 changes: 177 additions & 1 deletion tests/cppunit/types/tdigest_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -30,13 +30,14 @@
#include <range/v3/algorithm/shuffle.hpp>
#include <range/v3/range.hpp>
#include <range/v3/view/chunk.hpp>
#include <range/v3/view/concat.hpp>
#include <range/v3/view/iota.hpp>
#include <range/v3/view/join.hpp>
#include <range/v3/view/repeat.hpp>
#include <range/v3/view/transform.hpp>
#include <string>
#include <vector>

#include "logging.h"
#include "storage/redis_metadata.h"
#include "test_base.h"
#include "time_util.h"
Expand Down Expand Up @@ -948,3 +949,178 @@ TEST_F(RedisTDigestTest, MergeWithUserSpecifiedCompression) {
// Verify total observations: dest(1) + src(1) = 2
EXPECT_EQ(metadata.total_observations, 2);
}

TEST_F(RedisTDigestTest, CDFTest) {
std::string cdf_tdigest_name = "test_cdf_digest" + std::to_string(util::GetTimeStampMS());
bool exists = false;
auto status = tdigest_->Create(*ctx_, cdf_tdigest_name, {100}, &exists);
ASSERT_FALSE(exists);
ASSERT_TRUE(status.ok());

std::vector<double> samples = {1, 2, 2, 3, 3, 3, 4, 4, 4, 4, 5, 5, 5, 5, 5};
status = tdigest_->Add(*ctx_, cdf_tdigest_name, samples);
ASSERT_TRUE(status.ok());

std::vector<double> cdf_vals = {0, 1, 2, 3, 4, 5, 6};
redis::TDigestCDFResult result;

status = tdigest_->CDF(*ctx_, cdf_tdigest_name, cdf_vals, &result);
ASSERT_TRUE(status.ok()) << status.ToString();

std::vector<double> expected = {0.00, 0.03, 0.13, 0.29, 0.53, 0.83, 1.00};
ASSERT_EQ(result.cdf_values.size(), cdf_vals.size());

for (size_t i = 0; i < cdf_vals.size(); i++) {
EXPECT_NEAR(result.cdf_values[i], expected[i], 0.015) << fmt::format("Mismatch at index {}", i);
}
}

TEST_F(RedisTDigestTest, CDFReturnsNaNOnEmptyTDigest) {
std::string test_digest_name = "test_digest_cdf_nan" + std::to_string(util::GetTimeStampMS());

bool exists = false;
auto status = tdigest_->Create(*ctx_, test_digest_name, {100}, &exists);
ASSERT_FALSE(exists);
ASSERT_TRUE(status.ok());

std::vector<double> values = {0.0, 1.0, 2.0, 3.0};
redis::TDigestCDFResult result;

status = tdigest_->CDF(*ctx_, test_digest_name, values, &result);
ASSERT_TRUE(status.ok()) << status.ToString();
ASSERT_EQ(result.cdf_values.size(), values.size());
for (const auto cdf : result.cdf_values) {
EXPECT_TRUE(std::isnan(cdf));
}
}

TEST_F(RedisTDigestTest, CDFDuplicateValues) {
std::string test_digest_name = "test_cdf_duplicates" + std::to_string(util::GetTimeStampMS());

bool exists = false;
auto status = tdigest_->Create(*ctx_, test_digest_name, {100}, &exists);
ASSERT_FALSE(exists);
ASSERT_TRUE(status.ok());

status = tdigest_->Add(*ctx_, test_digest_name, {10, 10, 10, 20, 20});
ASSERT_TRUE(status.ok());

std::vector<double> cdf_vals = {5, 10, 20, 25};
redis::TDigestCDFResult result;
status = tdigest_->CDF(*ctx_, test_digest_name, cdf_vals, &result);
ASSERT_TRUE(status.ok()) << status.ToString();

std::vector<double> expected = {0, 0.3, 0.8, 1};
ASSERT_EQ(result.cdf_values.size(), cdf_vals.size());
for (size_t i = 0; i < cdf_vals.size(); i++) {
EXPECT_NEAR(result.cdf_values[i], expected[i], 0.001) << fmt::format("Mismatch at index {}", i);
}
}

TEST_F(RedisTDigestTest, CDFSignedZeroQueries) {
std::string test_digest_name = "test_cdf_signed_zero" + std::to_string(util::GetTimeStampMS());

bool exists = false;
auto status = tdigest_->Create(*ctx_, test_digest_name, {100}, &exists);
ASSERT_FALSE(exists);
ASSERT_TRUE(status.ok());

status = tdigest_->Add(*ctx_, test_digest_name, {-1, 0, 1});
ASSERT_TRUE(status.ok());

std::vector<double> cdf_vals = {-0.0, 0.0};
redis::TDigestCDFResult result;
status = tdigest_->CDF(*ctx_, test_digest_name, cdf_vals, &result);
ASSERT_TRUE(status.ok()) << status.ToString();

ASSERT_EQ(result.cdf_values.size(), cdf_vals.size());
EXPECT_NEAR(result.cdf_values[0], 0.5, 0.001);
EXPECT_NEAR(result.cdf_values[1], 0.5, 0.001);
}

TEST_F(RedisTDigestTest, CDFUniformDistribution) {
std::string test_digest_name = "test_cdf_uniform" + std::to_string(util::GetTimeStampMS());

bool exists = false;
auto status = tdigest_->Create(*ctx_, test_digest_name, {200}, &exists);
ASSERT_FALSE(exists);
ASSERT_TRUE(status.ok());

std::vector<double> samples = ranges::views::iota(1, 101) |
ranges::views::transform([](int i) { return (double)i; }) |
ranges::to<std::vector<double>>();
status = tdigest_->Add(*ctx_, test_digest_name, samples);
ASSERT_TRUE(status.ok());

std::vector<double> cdf_vals = {1, 25, 50, 75, 100};
redis::TDigestCDFResult result;
status = tdigest_->CDF(*ctx_, test_digest_name, cdf_vals, &result);
ASSERT_TRUE(status.ok()) << status.ToString();

std::vector<double> expected = {0.01, 0.25, 0.50, 0.75, 1.00};
ASSERT_EQ(result.cdf_values.size(), cdf_vals.size());

for (size_t i = 0; i < cdf_vals.size(); i++) {
EXPECT_NEAR(result.cdf_values[i], expected[i], 0.02) << fmt::format("Mismatch at index {}, val={}", i, cdf_vals[i]);
}
}

TEST_F(RedisTDigestTest, CDFMultipleAdds) {
std::string test_digest_name = "test_cdf_multiadd" + std::to_string(util::GetTimeStampMS());

bool exists = false;
auto status = tdigest_->Create(*ctx_, test_digest_name, {100}, &exists);
ASSERT_FALSE(exists);
ASSERT_TRUE(status.ok());

std::vector<double> samples1 = {1, 2, 3, 4, 5};
std::vector<double> samples2 = {6, 7, 8, 9, 10};
status = tdigest_->Add(*ctx_, test_digest_name, samples1);
ASSERT_TRUE(status.ok());
status = tdigest_->Add(*ctx_, test_digest_name, samples2);
ASSERT_TRUE(status.ok());

std::vector<double> cdf_vals = {1, 5, 7, 10};
redis::TDigestCDFResult result;
status = tdigest_->CDF(*ctx_, test_digest_name, cdf_vals, &result);
ASSERT_TRUE(status.ok());

std::vector<double> expected = {0.10, 0.50, 0.70, 1.00};
ASSERT_EQ(result.cdf_values.size(), cdf_vals.size());

for (size_t i = 0; i < cdf_vals.size(); i++) {
EXPECT_NEAR((result.cdf_values)[i], expected[i], 0.06)
<< fmt::format("Mismatch at index {}, val={}", i, cdf_vals[i]);
}
}

TEST_F(RedisTDigestTest, CDFSkewedDistribution) {
std::string test_digest_name = "test_cdf_skewed" + std::to_string(util::GetTimeStampMS());

bool exists = false;
auto status = tdigest_->Create(*ctx_, test_digest_name, {200}, &exists);
ASSERT_FALSE(exists);
ASSERT_TRUE(status.ok());

std::vector<double> samples =
ranges::views::concat(
ranges::views::repeat(0.0) | ranges::views::take(100),
ranges::views::iota(1, 11) | ranges::views::transform([](int i) { return static_cast<double>(i); })) |
ranges::to<std::vector<double>>();

status = tdigest_->Add(*ctx_, test_digest_name, samples);
ASSERT_TRUE(status.ok());

std::vector<double> cdf_vals = {0, 1, 5, 10};
redis::TDigestCDFResult result;
status = tdigest_->CDF(*ctx_, test_digest_name, cdf_vals, &result);
ASSERT_TRUE(status.ok());

std::vector<double> expected = {0.4545, 0.91, 0.95, 1.00};
ASSERT_EQ(result.cdf_values.size(), cdf_vals.size());

for (size_t i = 0; i < cdf_vals.size(); i++) {
EXPECT_NEAR((result.cdf_values)[i], expected[i], 0.03)
<< fmt::format("Mismatch at index {}, val={}", i, cdf_vals[i]);
}
}
Loading
Loading