This is an automated email from the ASF dual-hosted git repository.
tisonkun pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/datasketches-rust.git
The following commit(s) were added to refs/heads/main by this push:
new 7e578e7 perf(tdigest): reduce allocation overhead during compression
(#209)
7e578e7 is described below
commit 7e578e70e2c63fbcb2933a97c8ece6ffe604d0f8
Author: tison <[email protected]>
AuthorDate: Mon Aug 24 15:16:54 2026 +0800
perf(tdigest): reduce allocation overhead during compression (#209)
---
CONTRIBUTING.md | 6 +
Cargo.lock | 49 ++++
Cargo.toml | 1 +
datasketches/Cargo.toml | 7 +
datasketches/benches/tdigest.rs | 416 ++++++++++++++++++++++++++++++
datasketches/src/tdigest/sketch.rs | 90 +++++--
datasketches/tests/tdigest_test/sketch.rs | 10 +
xtask/src/main.rs | 18 ++
8 files changed, 569 insertions(+), 28 deletions(-)
diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md
index 9230f6a..48acafa 100644
--- a/CONTRIBUTING.md
+++ b/CONTRIBUTING.md
@@ -60,6 +60,12 @@ Lint:
cargo x lint
```
+Benchmark:
+
+```shell
+cargo x bench
+```
+
## Public API documentation
- Describe types with noun phrases and API behavior with third-person
present-tense verbs such as `Creates`, `Updates`, and `Returns`.
diff --git a/Cargo.lock b/Cargo.lock
index d18361b..c8ec060 100644
--- a/Cargo.lock
+++ b/Cargo.lock
@@ -165,6 +165,7 @@ dependencies = [
"anstyle",
"clap_lex",
"strsim",
+ "terminal_size",
]
[[package]]
@@ -191,6 +192,12 @@ version = "1.0.5"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "1d07550c9036bf2ae0c684c4297d503f838287c83c53686d05370d0e139ae570"
+[[package]]
+name = "condtype"
+version = "1.3.0"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "baf0a07a401f374238ab8e2f11a104d2851bf9ce711ec69804834de8af45c7af"
+
[[package]]
name = "console"
version = "0.16.4"
@@ -206,10 +213,36 @@ dependencies = [
name = "datasketches"
version = "0.4.0"
dependencies = [
+ "divan",
"googletest",
"insta",
]
+[[package]]
+name = "divan"
+version = "0.1.21"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "a405457ec78b8fe08b0e32b4a3570ab5dff6dd16eb9e76a5ee0a9d9cbd898933"
+dependencies = [
+ "cfg-if",
+ "clap",
+ "condtype",
+ "divan-macros",
+ "libc",
+ "regex-lite",
+]
+
+[[package]]
+name = "divan-macros"
+version = "0.1.21"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "9556bc800956545d6420a640173e5ba7dfa82f38d3ea5a167eb555bc69ac3323"
+dependencies = [
+ "proc-macro2",
+ "quote",
+ "syn 2.0.119",
+]
+
[[package]]
name = "encode_unicode"
version = "1.0.0"
@@ -457,6 +490,12 @@ dependencies = [
"regex-syntax",
]
+[[package]]
+name = "regex-lite"
+version = "0.1.9"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "cab834c73d247e67f4fae452806d17d3c7501756d98c8808d7c9c7aa7d18f973"
+
[[package]]
name = "regex-syntax"
version = "0.8.11"
@@ -670,6 +709,16 @@ dependencies = [
"windows-sys 0.61.2",
]
+[[package]]
+name = "terminal_size"
+version = "0.4.4"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "230a1b821ccbd75b185820a1f1ff7b14d21da1e442e22c0863ea5f08771a8874"
+dependencies = [
+ "rustix",
+ "windows-sys 0.61.2",
+]
+
[[package]]
name = "thiserror"
version = "2.0.19"
diff --git a/Cargo.toml b/Cargo.toml
index 04f944e..36e2cea 100644
--- a/Cargo.toml
+++ b/Cargo.toml
@@ -33,6 +33,7 @@ datasketches = { path = "datasketches" }
# Crates.io dependencies
clap = { version = "4.6.5", features = ["derive"] }
+divan = { version = "0.1.21" }
insta = { version = "1.48.0" }
googletest = { version = "0.14.3" }
cargo_metadata = { version = "0.23.1" }
diff --git a/datasketches/Cargo.toml b/datasketches/Cargo.toml
index 169ee7e..8ff9adb 100644
--- a/datasketches/Cargo.toml
+++ b/datasketches/Cargo.toml
@@ -87,7 +87,14 @@ name = "tuple_test"
path = "tests/tuple_test/main.rs"
required-features = ["tuple"]
+[[bench]]
+name = "tdigest"
+path = "benches/tdigest.rs"
+required-features = ["tdigest"]
+harness = false
+
[dev-dependencies]
+divan = { workspace = true }
googletest = { workspace = true }
insta = { workspace = true }
diff --git a/datasketches/benches/tdigest.rs b/datasketches/benches/tdigest.rs
new file mode 100644
index 0000000..b68bd59
--- /dev/null
+++ b/datasketches/benches/tdigest.rs
@@ -0,0 +1,416 @@
+// Licensed to the Apache Software Foundation (ASF) under one
+// or more contributor license agreements. See the NOTICE file
+// distributed with this work for additional information
+// regarding copyright ownership. The ASF licenses this file
+// to you under the Apache License, Version 2.0 (the
+// "License"); you may not use this file except in compliance
+// with the License. You may obtain a copy of the License at
+//
+// http://www.apache.org/licenses/LICENSE-2.0
+//
+// Unless required by applicable law or agreed to in writing,
+// software distributed under the License is distributed on an
+// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+// KIND, either express or implied. See the License for the
+// specific language governing permissions and limitations
+// under the License.
+
+use datasketches::tdigest::TDigest;
+use datasketches::tdigest::TDigestMut;
+use divan::AllocProfiler;
+use divan::Bencher;
+use divan::black_box;
+use divan::black_box_drop;
+use divan::counter::BytesCount;
+use divan::counter::ItemsCount;
+
+#[global_allocator]
+static ALLOC: AllocProfiler = AllocProfiler::system();
+
+const PARTIAL_GROUPS: usize = 512;
+const SMALL_ROWS_PER_PARTIAL: usize = 8;
+const DEFAULT_DIGEST_K: u16 = 200;
+const ROWS_PER_PARTIAL: usize = 64;
+
+fn main() {
+ divan::main();
+}
+
+#[divan::bench(args = [1_000, 100_000])]
+fn update(bencher: Bencher, len: usize) {
+ let values = values(len);
+
+ bencher
+ .counter(ItemsCount::new(len))
+ .bench_local(|| build_digest(black_box(&values)));
+}
+
+#[divan::bench(args = [1, 2, 8, 32, 64, 128])]
+fn small_digest_lifecycle(bencher: Bencher, rows: usize) {
+ let values = values(rows);
+ let (serialized_bytes, _) = serialized_state_shape(200, &values);
+
+ bencher
+ .counter(ItemsCount::new(rows))
+ .counter(BytesCount::new(serialized_bytes))
+ .bench_local(|| {
+ let mut digest = TDigestMut::default();
+ for &value in black_box(&values) {
+ digest.update(value);
+ }
+ let bytes = digest.serialize();
+ black_box_drop(bytes);
+ black_box_drop(digest);
+ });
+}
+
+#[divan::bench(args = [10_u16, 200_u16])]
+fn partial_digest_lifecycle_by_k(bencher: Bencher, k: u16) {
+ let values = values(ROWS_PER_PARTIAL);
+ let (serialized_bytes, centroids) = serialized_state_shape(k, &values);
+ assert!(matches!(
+ (k, serialized_bytes, centroids),
+ (10, 224, 12) | (200, 1_056, 64)
+ ));
+
+ bencher
+ .counter(ItemsCount::new(ROWS_PER_PARTIAL))
+ .counter(BytesCount::new(serialized_bytes))
+ .bench_local(|| {
+ let mut digest = TDigestMut::new(black_box(k));
+ for &value in black_box(&values) {
+ digest.update(value);
+ }
+ let bytes = digest.serialize();
+ black_box_drop(bytes);
+ black_box_drop(digest);
+ });
+}
+
+#[divan::bench]
+fn compress(bencher: Bencher) {
+ // The default k=200 digest buffers 1,640 values before automatic
compression.
+ let values = values(1_640);
+ let digest = build_mut_digest(&values);
+
+ bencher
+ .counter(ItemsCount::new(values.len()))
+ .with_inputs(|| digest.clone())
+ .bench_local_values(|digest| black_box(digest.freeze()));
+}
+
+#[divan::bench]
+fn merge(bencher: Bencher) {
+ let values = values(200_000);
+ let mut left = build_mut_digest(&values[..100_000]);
+ let mut right = build_mut_digest(&values[100_000..]);
+ black_box(left.rank(0.5));
+ black_box(right.rank(0.5));
+
+ bencher
+ .counter(ItemsCount::new(values.len()))
+ .with_inputs(|| left.clone())
+ .bench_local_values(|mut left| {
+ left.merge(black_box(&right));
+ black_box(left)
+ });
+}
+
+#[divan::bench]
+fn rank(bencher: Bencher) {
+ let digest = prepared_digest();
+
+ bencher.bench_local(|| black_box(&digest).rank(black_box(0.531_25)));
+}
+
+#[divan::bench]
+fn quantile(bencher: Bencher) {
+ let digest = prepared_digest();
+
+ bencher.bench_local(|| black_box(&digest).quantile(black_box(0.531_25)));
+}
+
+#[divan::bench]
+fn cdf_100(bencher: Bencher) {
+ let digest = prepared_digest();
+ let split_points = (1..=100).map(|i| i as f64 / 101.0).collect::<Vec<_>>();
+
+ bencher
+ .counter(ItemsCount::new(split_points.len()))
+ .bench_local(|| black_box(&digest).cdf(black_box(&split_points)));
+}
+
+#[divan::bench]
+fn quantiles_2_sequential(bencher: Bencher) {
+ let digest = prepared_digest();
+
+ bencher.bench_local(|| {
+ [
+ black_box(&digest).quantile(black_box(0.5)),
+ black_box(&digest).quantile(black_box(0.95)),
+ ]
+ });
+}
+
+#[divan::bench]
+fn small_partial_groups_update(bencher: Bencher) {
+ bencher
+ .counter(ItemsCount::new(PARTIAL_GROUPS * SMALL_ROWS_PER_PARTIAL))
+ .bench_local(|| {
+ let mut digests = (0..PARTIAL_GROUPS)
+ .map(|_| TDigestMut::default())
+ .collect::<Vec<_>>();
+ for (group, digest) in digests.iter_mut().enumerate() {
+ for row in 0..SMALL_ROWS_PER_PARTIAL {
+ digest.update(partial_value(group, row,
SMALL_ROWS_PER_PARTIAL));
+ }
+ }
+ black_box(digests)
+ });
+}
+
+#[divan::bench]
+fn small_partial_groups_update_two_states(bencher: Bencher) {
+ bencher
+ .counter(ItemsCount::new(PARTIAL_GROUPS * SMALL_ROWS_PER_PARTIAL))
+ .bench_local(|| {
+ let mut digests = (0..PARTIAL_GROUPS)
+ .map(|_| (TDigestMut::default(), TDigestMut::default()))
+ .collect::<Vec<_>>();
+ for (group, (first, second)) in digests.iter_mut().enumerate() {
+ for row in 0..SMALL_ROWS_PER_PARTIAL {
+ let value = partial_value(group, row,
SMALL_ROWS_PER_PARTIAL);
+ first.update(value);
+ second.update(value);
+ }
+ }
+ black_box(digests)
+ });
+}
+
+#[divan::bench]
+fn small_partial_groups_serialize(bencher: Bencher) {
+ let digests = partial_digests(PARTIAL_GROUPS, SMALL_ROWS_PER_PARTIAL);
+
+ bencher
+ .counter(ItemsCount::new(PARTIAL_GROUPS))
+ .with_inputs(|| digests.clone())
+ .bench_local_values(|mut digests| {
+ let bytes = digests
+ .iter_mut()
+ .map(TDigestMut::serialize)
+ .collect::<Vec<_>>();
+ black_box(bytes)
+ });
+}
+
+#[divan::bench]
+fn small_partial_groups_deserialize(bencher: Bencher) {
+ let bytes = serialized_partial_digests(PARTIAL_GROUPS,
SMALL_ROWS_PER_PARTIAL);
+
+ bencher
+ .counter(ItemsCount::new(PARTIAL_GROUPS))
+ .bench_local(|| {
+ let digests = bytes
+ .iter()
+ .map(|bytes| TDigestMut::deserialize(bytes, false).unwrap())
+ .collect::<Vec<_>>();
+ black_box(digests)
+ });
+}
+
+#[divan::bench]
+fn small_partial_merge(bencher: Bencher) {
+ let partials = partial_digests(64, SMALL_ROWS_PER_PARTIAL)
+ .into_iter()
+ .map(|mut digest| {
+ black_box(digest.rank(0.0));
+ digest
+ })
+ .collect::<Vec<_>>();
+
+ bencher
+ .counter(ItemsCount::new(64 * SMALL_ROWS_PER_PARTIAL))
+ .bench_local(|| {
+ let mut merged = TDigestMut::default();
+ for partial in &partials {
+ merged.merge(black_box(partial));
+ }
+ black_box(merged)
+ });
+}
+
+#[divan::bench]
+fn partial_groups_update(bencher: Bencher) {
+ bencher
+ .counter(ItemsCount::new(PARTIAL_GROUPS * ROWS_PER_PARTIAL))
+ .bench_local(|| {
+ let mut digests = (0..PARTIAL_GROUPS)
+ .map(|_| TDigestMut::new(DEFAULT_DIGEST_K))
+ .collect::<Vec<_>>();
+ for (group, digest) in digests.iter_mut().enumerate() {
+ for row in 0..ROWS_PER_PARTIAL {
+ digest.update(partial_value(group, row, ROWS_PER_PARTIAL));
+ }
+ }
+ black_box(digests)
+ });
+}
+
+#[divan::bench]
+fn partial_groups_update_two_states(bencher: Bencher) {
+ bencher
+ .counter(ItemsCount::new(PARTIAL_GROUPS * ROWS_PER_PARTIAL))
+ .bench_local(|| {
+ let mut digests = (0..PARTIAL_GROUPS)
+ .map(|_| {
+ (
+ TDigestMut::new(DEFAULT_DIGEST_K),
+ TDigestMut::new(DEFAULT_DIGEST_K),
+ )
+ })
+ .collect::<Vec<_>>();
+ for (group, (first, second)) in digests.iter_mut().enumerate() {
+ for row in 0..ROWS_PER_PARTIAL {
+ let value = partial_value(group, row, ROWS_PER_PARTIAL);
+ first.update(value);
+ second.update(value);
+ }
+ }
+ black_box(digests)
+ });
+}
+
+#[divan::bench]
+fn partial_groups_serialize(bencher: Bencher) {
+ let digests = partial_digests_with(DEFAULT_DIGEST_K, PARTIAL_GROUPS,
ROWS_PER_PARTIAL);
+
+ bencher
+ .counter(ItemsCount::new(PARTIAL_GROUPS))
+ .with_inputs(|| digests.clone())
+ .bench_local_values(|mut digests| {
+ let bytes = digests
+ .iter_mut()
+ .map(TDigestMut::serialize)
+ .collect::<Vec<_>>();
+ black_box(bytes)
+ });
+}
+
+#[divan::bench]
+fn partial_groups_deserialize(bencher: Bencher) {
+ let bytes = serialized_partial_digests_with(DEFAULT_DIGEST_K,
PARTIAL_GROUPS, ROWS_PER_PARTIAL);
+
+ bencher
+ .counter(ItemsCount::new(PARTIAL_GROUPS))
+ .bench_local(|| {
+ let digests = bytes
+ .iter()
+ .map(|bytes| TDigestMut::deserialize(bytes, false).unwrap())
+ .collect::<Vec<_>>();
+ black_box(digests)
+ });
+}
+
+#[divan::bench]
+fn partial_merge(bencher: Bencher) {
+ let partials = partial_digests_with(DEFAULT_DIGEST_K, 64, ROWS_PER_PARTIAL)
+ .into_iter()
+ .map(|mut digest| {
+ black_box(digest.rank(0.0));
+ digest
+ })
+ .collect::<Vec<_>>();
+
+ bencher
+ .counter(ItemsCount::new(64 * ROWS_PER_PARTIAL))
+ .bench_local(|| {
+ let mut merged = TDigestMut::new(DEFAULT_DIGEST_K);
+ for partial in &partials {
+ merged.merge(black_box(partial));
+ }
+ black_box(merged)
+ });
+}
+
+fn prepared_digest() -> TDigest {
+ build_digest(&values(100_000))
+}
+
+fn build_digest(values: &[f64]) -> TDigest {
+ build_mut_digest(values).freeze()
+}
+
+fn build_mut_digest(values: &[f64]) -> TDigestMut {
+ let mut digest = TDigestMut::default();
+ for &value in values {
+ digest.update(value);
+ }
+ digest
+}
+
+fn partial_digests(groups: usize, rows_per_group: usize) -> Vec<TDigestMut> {
+ partial_digests_with(200, groups, rows_per_group)
+}
+
+fn partial_digests_with(k: u16, groups: usize, rows_per_group: usize) ->
Vec<TDigestMut> {
+ (0..groups)
+ .map(|group| {
+ let mut digest = TDigestMut::new(k);
+ for row in 0..rows_per_group {
+ digest.update(partial_value(group, row, rows_per_group));
+ }
+ digest
+ })
+ .collect()
+}
+
+fn serialized_partial_digests(groups: usize, rows_per_group: usize) ->
Vec<Vec<u8>> {
+ let bytes = serialized_partial_digests_with(200, groups, rows_per_group);
+ assert!(
+ bytes
+ .iter()
+ .all(|bytes| bytes.len() == 32 + rows_per_group * 16)
+ );
+ assert!(bytes.iter().all(|bytes| {
+ u32::from_le_bytes(bytes[8..12].try_into().unwrap()) == rows_per_group
as u32
+ }));
+ bytes
+}
+
+fn serialized_partial_digests_with(k: u16, groups: usize, rows_per_group:
usize) -> Vec<Vec<u8>> {
+ partial_digests_with(k, groups, rows_per_group)
+ .into_iter()
+ .map(|mut digest| digest.serialize())
+ .collect()
+}
+
+fn serialized_state_shape(k: u16, values: &[f64]) -> (usize, u32) {
+ let mut digest = TDigestMut::new(k);
+ for &value in values {
+ digest.update(value);
+ }
+ let bytes = digest.serialize();
+ let centroids = match values.len() {
+ 0 => 0,
+ 1 => 1,
+ _ => u32::from_le_bytes(bytes[8..12].try_into().unwrap()),
+ };
+ (bytes.len(), centroids)
+}
+
+fn partial_value(group: usize, row: usize, rows_per_group: usize) -> f64 {
+ (group * rows_per_group + row) as f64
+}
+
+fn values(len: usize) -> Vec<f64> {
+ let mut state = 0x9e37_79b9_7f4a_7c15_u64;
+ (0..len)
+ .map(|_| {
+ state ^= state << 13;
+ state ^= state >> 7;
+ state ^= state << 17;
+ (state >> 11) as f64 * (1.0 / ((1_u64 << 53) as f64))
+ })
+ .collect()
+}
diff --git a/datasketches/src/tdigest/sketch.rs
b/datasketches/src/tdigest/sketch.rs
index 8ae557f..7e216f3 100644
--- a/datasketches/src/tdigest/sketch.rs
+++ b/datasketches/src/tdigest/sketch.rs
@@ -39,6 +39,8 @@ use crate::tdigest::serialization::SERIAL_VERSION;
const DEFAULT_K: u16 = 200;
/// Multiplier for buffer size relative to centroids capacity.
const BUFFER_MULTIPLIER: usize = 4;
+/// Buffer capacity allocated by the first update to a digest.
+const INITIAL_BUFFER_CAPACITY: usize = 8;
/// Default weight for single values.
const DEFAULT_WEIGHT: NonZeroU64 = NonZeroU64::new(1).unwrap();
@@ -134,18 +136,15 @@ impl TDigestMut {
reverse_merge: bool,
min: f64,
max: f64,
- mut centroids: Vec<Centroid>,
+ centroids: Vec<Centroid>,
centroids_weight: u64,
- mut buffer: Vec<f64>,
+ buffer: Vec<f64>,
) -> Self {
assert!(k >= 10, "k must be at least 10");
let fudge = if k < 30 { 30 } else { 10 };
let centroids_capacity = (k as usize * 2) + fudge;
- centroids.reserve(centroids_capacity);
- buffer.reserve(centroids_capacity * BUFFER_MULTIPLIER);
-
TDigestMut {
k,
reverse_merge,
@@ -172,13 +171,31 @@ impl TDigestMut {
/// assert!(sketch.total_weight() >= 1);
/// ```
pub fn update(&mut self, value: f64) {
- if value.is_nan() || value.is_infinite() {
+ if !value.is_finite() {
return;
}
- if self.buffer.len() == self.centroids_capacity * BUFFER_MULTIPLIER {
+ let full_buffer_capacity = self.centroids_capacity * BUFFER_MULTIPLIER;
+ if self.buffer.len() == full_buffer_capacity {
self.compress();
}
+ if self.buffer.len() == self.buffer.capacity() {
+ let target_capacity = if self.buffer.capacity() == 0 {
+ INITIAL_BUFFER_CAPACITY
+ } else if self.buffer.capacity() == INITIAL_BUFFER_CAPACITY {
+ // Once a digest outgrows a tiny group, skip an extra
allocator round trip while
+ // keeping the first allocation small.
+ (INITIAL_BUFFER_CAPACITY * BUFFER_MULTIPLIER *
BUFFER_MULTIPLIER)
+ .min(full_buffer_capacity)
+ } else {
+ self.buffer
+ .capacity()
+ .saturating_mul(BUFFER_MULTIPLIER)
+ .min(full_buffer_capacity)
+ };
+ self.buffer
+
.reserve_exact(target_capacity.saturating_sub(self.buffer.len()));
+ }
self.buffer.push(value);
self.min = self.min.min(value);
@@ -237,9 +254,10 @@ impl TDigestMut {
return;
}
- let mut tmp = Vec::with_capacity(
- self.centroids.len() + self.buffer.len() + other.centroids.len() +
other.buffer.len(),
- );
+ let additional_centroids = self.buffer.len() + other.centroids.len() +
other.buffer.len();
+ let mut tmp = std::mem::take(&mut self.centroids);
+ let num_existing_centroids = tmp.len();
+ tmp.reserve(additional_centroids);
for &v in &self.buffer {
tmp.push(Centroid {
mean: v,
@@ -255,6 +273,9 @@ impl TDigestMut {
for &c in &other.centroids {
tmp.push(c);
}
+ // Preserve the original insertion order for equal means because
t-digest requires a stable
+ // sort: buffered values, the other digest's centroids, then this
digest's centroids.
+ tmp.rotate_left(num_existing_centroids);
self.do_merge(tmp, self.buffer.len() as u64 + other.total_weight())
}
@@ -747,13 +768,25 @@ impl TDigestMut {
if self.buffer.is_empty() {
return;
}
- let mut tmp = Vec::with_capacity(self.buffer.len() +
self.centroids.len());
+ let mut tmp = std::mem::take(&mut self.centroids);
+ let full_buffer_capacity = self.centroids_capacity * BUFFER_MULTIPLIER;
+ let required_capacity = tmp.len() + self.buffer.len();
+ let target_capacity = if self.buffer.len() == full_buffer_capacity {
+ required_capacity.max(full_buffer_capacity +
self.centroids_capacity)
+ } else {
+ required_capacity
+ };
+ tmp.reserve_exact(target_capacity.saturating_sub(tmp.len()));
+ let num_buffered = self.buffer.len();
for &v in &self.buffer {
tmp.push(Centroid {
mean: v,
weight: DEFAULT_WEIGHT,
});
}
+ // Preserve the original insertion order for equal means because
t-digest requires a stable
+ // sort: buffered values before existing centroids.
+ tmp.rotate_right(num_buffered);
self.do_merge(tmp, self.buffer.len() as u64)
}
@@ -761,33 +794,32 @@ impl TDigestMut {
///
/// # Contract
///
- /// * `buffer` must have at least one centroid.
- /// * `buffer` is generated from `self.buffer`, and thus:
- /// * No `NAN` values are present in `buffer`.
- /// * We should clear `self.buffer` after merging.
+ /// * `buffer` must contain at least one centroid.
+ /// * `buffer` contains every centroid to be merged, including all
centroids previously stored
+ /// in `self`.
+ /// * `weight` is the total weight not yet included in
`self.centroids_weight`.
+ /// * Every centroid mean in `buffer` is finite.
+ /// * `self.buffer` is cleared before returning.
fn do_merge(&mut self, mut buffer: Vec<Centroid>, weight: u64) {
- buffer.extend(std::mem::take(&mut self.centroids));
buffer.sort_by(centroid_cmp);
if self.reverse_merge {
buffer.reverse();
}
self.centroids_weight += weight;
- let mut num_centroids = 0;
+ let mut num_centroids = 1;
let len = buffer.len();
- self.centroids.push(buffer[0]);
- num_centroids += 1;
+ let centroids_weight = self.centroids_weight as f64;
+ let normalizer = scale_function::normalizer(2.0 * f64::from(self.k),
centroids_weight);
let mut current = 1;
let mut weight_so_far = 0.;
while current < len {
let c = buffer[current];
- let proposed_weight = self.centroids[num_centroids - 1].weight() +
c.weight();
+ let proposed_weight = buffer[num_centroids - 1].weight() +
c.weight();
let mut add_this = false;
if (current != 1) && (current != (len - 1)) {
- let centroids_weight = self.centroids_weight as f64;
let q0 = weight_so_far / centroids_weight;
let q2 = (weight_so_far + proposed_weight) / centroids_weight;
- let normalizer = scale_function::normalizer((2 * self.k) as
f64, centroids_weight);
add_this = proposed_weight
<= (centroids_weight
* scale_function::max(q0, normalizer)
@@ -795,21 +827,23 @@ impl TDigestMut {
}
if add_this {
// merge into existing centroid
- self.centroids[num_centroids - 1].add(c);
+ buffer[num_centroids - 1].add(c);
} else {
// copy to a new centroid
- weight_so_far += self.centroids[num_centroids - 1].weight();
- self.centroids.push(c);
+ weight_so_far += buffer[num_centroids - 1].weight();
+ buffer[num_centroids] = c;
num_centroids += 1;
}
current += 1;
}
+ buffer.truncate(num_centroids);
if self.reverse_merge {
- self.centroids.reverse();
+ buffer.reverse();
}
- self.min = self.min.min(self.centroids[0].mean);
- self.max = self.max.max(self.centroids[num_centroids - 1].mean);
+ self.min = self.min.min(buffer[0].mean);
+ self.max = self.max.max(buffer[num_centroids - 1].mean);
+ self.centroids = buffer;
self.reverse_merge = !self.reverse_merge;
self.buffer.clear();
}
diff --git a/datasketches/tests/tdigest_test/sketch.rs
b/datasketches/tests/tdigest_test/sketch.rs
index 13ca646..3094194 100644
--- a/datasketches/tests/tdigest_test/sketch.rs
+++ b/datasketches/tests/tdigest_test/sketch.rs
@@ -65,6 +65,16 @@ fn test_one_value() {
assert_eq!(tdigest.quantile(1.0), Some(1.0));
}
+#[test]
+fn test_maximum_k() {
+ let mut tdigest = TDigestMut::new(u16::MAX);
+ tdigest.update(1.0);
+
+ let tdigest = tdigest.freeze();
+ assert_eq!(tdigest.k(), u16::MAX);
+ assert_eq!(tdigest.quantile(0.5), Some(1.0));
+}
+
#[test]
fn test_many_values() {
let n = 10000;
diff --git a/xtask/src/main.rs b/xtask/src/main.rs
index dc4fe43..bac79cd 100644
--- a/xtask/src/main.rs
+++ b/xtask/src/main.rs
@@ -43,6 +43,7 @@ struct Command {
impl Command {
fn run(self) {
match self.sub {
+ SubCommand::Bench(cmd) => cmd.run(),
SubCommand::Check(cmd) => cmd.run(),
SubCommand::Docs(cmd) => cmd.run(),
SubCommand::Lint(cmd) => cmd.run(),
@@ -54,6 +55,8 @@ impl Command {
#[derive(Subcommand)]
enum SubCommand {
+ #[clap(about = "Run workspace benchmarks.")]
+ Bench(CommandBench),
#[clap(about = "Check datasketches under the feature matrix.")]
Check(CommandCheck),
#[clap(about = "Generate documentation and open for preview")]
@@ -69,6 +72,15 @@ enum SubCommand {
PrepareTestData(CommandPrepareTestData),
}
+#[derive(Parser)]
+struct CommandBench;
+
+impl CommandBench {
+ fn run(self) {
+ run_command(make_bench_cmd());
+ }
+}
+
#[derive(Parser)]
#[clap(name = "check")]
struct CommandCheck {}
@@ -178,6 +190,12 @@ fn run_command(mut cmd: StdCommand) {
assert!(status.success(), "command failed: {status}");
}
+fn make_bench_cmd() -> StdCommand {
+ let mut cmd = find_command("cargo");
+ cmd.args(["bench", "--workspace", "--all-features", "--bench", "*"]);
+ cmd
+}
+
fn make_test_cmd(no_capture: bool, features: &[String]) -> StdCommand {
let mut cmd = find_command("cargo");
cmd.args(["test", "--workspace", "--no-default-features"]);
---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]