Skip to content

Commit fae8080

Browse files
authored
[core] Optimize TDigest percentile aggregation (#18996)
1 parent dce035f commit fae8080

15 files changed

Lines changed: 4251 additions & 60 deletions

File tree

pinot-core/src/main/java/org/apache/pinot/core/query/aggregation/function/PercentileTDigestAccumulator.java

Lines changed: 1049 additions & 0 deletions
Large diffs are not rendered by default.

pinot-core/src/main/java/org/apache/pinot/core/query/aggregation/function/PercentileTDigestAggregationFunction.java

Lines changed: 80 additions & 50 deletions
Original file line numberDiff line numberDiff line change
@@ -108,20 +108,18 @@ public void aggregate(int length, AggregationResultHolder aggregationResultHolde
108108
if (blockValSet.getValueType() == DataType.BYTES) {
109109
// Serialized TDigest
110110
byte[][] bytesValues = blockValSet.getBytesValuesSV();
111-
foldNotNull(length, blockValSet, (TDigest) aggregationResultHolder.getResult(), (tDigest, from, toEx) -> {
112-
if (tDigest != null) {
113-
for (int i = from; i < toEx; i++) {
114-
tDigest.add(ObjectSerDeUtils.TDIGEST_SER_DE.deserialize(bytesValues[i]));
115-
}
116-
} else {
117-
tDigest = ObjectSerDeUtils.TDIGEST_SER_DE.deserialize(bytesValues[0]);
118-
aggregationResultHolder.setValue(tDigest);
119-
for (int i = 1; i < length; i++) {
120-
tDigest.add(ObjectSerDeUtils.TDIGEST_SER_DE.deserialize(bytesValues[i]));
111+
foldNotNull(length, blockValSet,
112+
(PercentileTDigestAccumulator) aggregationResultHolder.getResult(), (current, from, toExclusive) -> {
113+
if (current == null) {
114+
current = PercentileTDigestAccumulator.forSerializedTDigestWithMergeBuffers(bytesValues[from]);
115+
aggregationResultHolder.setValue(current);
116+
}
117+
for (int i = from; i < toExclusive; i++) {
118+
current.addSerializedTDigest(bytesValues[i]);
119+
}
120+
return current;
121121
}
122-
}
123-
return tDigest;
124-
});
122+
);
125123
return;
126124
}
127125

@@ -134,22 +132,17 @@ public void aggregate(int length, AggregationResultHolder aggregationResultHolde
134132

135133
protected void aggregateSV(int length, AggregationResultHolder aggregationResultHolder, BlockValSet blockValSet) {
136134
double[] doubleValues = blockValSet.getDoubleValuesSV();
137-
TDigest tDigest = getDefaultTDigest(aggregationResultHolder, _compressionFactor);
138-
forEachNotNull(length, blockValSet, (from, to) -> {
139-
for (int i = from; i < to; i++) {
140-
tDigest.add(doubleValues[i]);
141-
}
142-
});
135+
PercentileTDigestAccumulator accumulator = getAccumulator(aggregationResultHolder, _compressionFactor);
136+
forEachNotNull(length, blockValSet, (from, to) -> accumulator.add(doubleValues, from, to));
143137
}
144138

145139
protected void aggregateMV(int length, AggregationResultHolder aggregationResultHolder, BlockValSet blockValSet) {
146140
double[][] valuesArray = blockValSet.getDoubleValuesMV();
147-
TDigest tDigest = getDefaultTDigest(aggregationResultHolder, _compressionFactor);
141+
PercentileTDigestAccumulator accumulator = getAccumulator(aggregationResultHolder, _compressionFactor);
148142
forEachNotNull(length, blockValSet, (from, to) -> {
149143
for (int i = from; i < to; i++) {
150-
for (double value : valuesArray[i]) {
151-
tDigest.add(value);
152-
}
144+
double[] values = valuesArray[i];
145+
accumulator.add(values, 0, values.length);
153146
}
154147
});
155148
}
@@ -163,14 +156,9 @@ public void aggregateGroupBySV(int length, int[] groupKeyArray, GroupByResultHol
163156
byte[][] bytesValues = blockValSet.getBytesValuesSV();
164157
forEachNotNull(length, blockValSet, (from, to) -> {
165158
for (int i = from; i < to; i++) {
166-
TDigest value = ObjectSerDeUtils.TDIGEST_SER_DE.deserialize(bytesValues[i]);
167159
int groupKey = groupKeyArray[i];
168-
TDigest tDigest = groupByResultHolder.getResult(groupKey);
169-
if (tDigest != null) {
170-
tDigest.add(value);
171-
} else {
172-
groupByResultHolder.setValueForKey(groupKey, value);
173-
}
160+
getSerializedAccumulator(groupByResultHolder, groupKey, bytesValues[i])
161+
.addSerializedTDigest(bytesValues[i]);
174162
}
175163
});
176164
return;
@@ -213,17 +201,17 @@ public void aggregateGroupByMV(int length, int[][] groupKeysArray, GroupByResult
213201
if (blockValSet.getValueType() == DataType.BYTES) {
214202
// Serialized TDigest
215203
byte[][] bytesValues = blockValSet.getBytesValuesSV();
204+
PercentileTDigestAccumulator.SerializedTDigestInput input =
205+
new PercentileTDigestAccumulator.SerializedTDigestInput();
216206
forEachNotNull(length, blockValSet, (from, to) -> {
217207
for (int i = from; i < to; i++) {
218-
TDigest value = ObjectSerDeUtils.TDIGEST_SER_DE.deserialize(bytesValues[i]);
219-
for (int groupKey : groupKeysArray[i]) {
220-
TDigest tDigest = groupByResultHolder.getResult(groupKey);
221-
if (tDigest != null) {
222-
tDigest.add(value);
223-
} else {
224-
// Create a new TDigest for the group
225-
groupByResultHolder.setValueForKey(groupKey, ObjectSerDeUtils.TDIGEST_SER_DE.deserialize(bytesValues[i]));
226-
}
208+
int[] groupKeys = groupKeysArray[i];
209+
if (groupKeys.length == 0) {
210+
continue;
211+
}
212+
input.reset(bytesValues[i]);
213+
for (int groupKey : groupKeys) {
214+
getSerializedAccumulator(groupByResultHolder, groupKey, input).addSerializedTDigest(input);
227215
}
228216
}
229217
});
@@ -268,21 +256,21 @@ protected void aggregateMVGroupByMV(int length, int[][] groupKeysArray, GroupByR
268256

269257
@Override
270258
public TDigest extractAggregationResult(AggregationResultHolder aggregationResultHolder) {
271-
TDigest tDigest = aggregationResultHolder.getResult();
272-
if (tDigest == null) {
259+
PercentileTDigestAccumulator accumulator = aggregationResultHolder.getResult();
260+
if (accumulator == null) {
273261
return TDigest.createMergingDigest(_compressionFactor);
274262
} else {
275-
return tDigest;
263+
return accumulator;
276264
}
277265
}
278266

279267
@Override
280268
public TDigest extractGroupByResult(GroupByResultHolder groupByResultHolder, int groupKey) {
281-
TDigest tDigest = groupByResultHolder.getResult(groupKey);
282-
if (tDigest == null) {
269+
Object result = groupByResultHolder.getResult(groupKey);
270+
if (result == null) {
283271
return TDigest.createMergingDigest(_compressionFactor);
284272
} else {
285-
return tDigest;
273+
return (TDigest) result;
286274
}
287275
}
288276

@@ -294,8 +282,15 @@ public TDigest merge(TDigest intermediateResult1, TDigest intermediateResult2) {
294282
if (intermediateResult2.size() == 0L) {
295283
return intermediateResult1;
296284
}
297-
intermediateResult1.add(intermediateResult2);
298-
return intermediateResult1;
285+
if (intermediateResult1 instanceof PercentileTDigestAccumulator) {
286+
intermediateResult1.add(intermediateResult2);
287+
return intermediateResult1;
288+
}
289+
PercentileTDigestAccumulator accumulator =
290+
PercentileTDigestAccumulator.forReduction(intermediateResult1.compression());
291+
accumulator.add(intermediateResult1);
292+
accumulator.add(intermediateResult2);
293+
return accumulator;
299294
}
300295

301296
@Override
@@ -305,13 +300,18 @@ public ColumnDataType getIntermediateResultColumnType() {
305300

306301
@Override
307302
public SerializedIntermediateResult serializeIntermediateResult(TDigest tDigest) {
308-
return new SerializedIntermediateResult(ObjectSerDeUtils.ObjectType.TDigest.getValue(),
309-
ObjectSerDeUtils.TDIGEST_SER_DE.serialize(tDigest));
303+
byte[] bytes;
304+
if (tDigest instanceof PercentileTDigestAccumulator) {
305+
bytes = ((PercentileTDigestAccumulator) tDigest).serialize();
306+
} else {
307+
bytes = ObjectSerDeUtils.TDIGEST_SER_DE.serialize(tDigest);
308+
}
309+
return new SerializedIntermediateResult(ObjectSerDeUtils.ObjectType.TDigest.getValue(), bytes);
310310
}
311311

312312
@Override
313313
public TDigest deserializeIntermediateResult(CustomObject customObject) {
314-
return ObjectSerDeUtils.TDIGEST_SER_DE.deserialize(customObject.getBuffer());
314+
return PercentileTDigestAccumulator.forSerializedTDigest(customObject.getBuffer());
315315
}
316316

317317
@Override
@@ -353,6 +353,36 @@ protected static TDigest getDefaultTDigest(AggregationResultHolder aggregationRe
353353
return tDigest;
354354
}
355355

356+
private static PercentileTDigestAccumulator getAccumulator(AggregationResultHolder aggregationResultHolder,
357+
int compressionFactor) {
358+
PercentileTDigestAccumulator accumulator = aggregationResultHolder.getResult();
359+
if (accumulator == null) {
360+
accumulator = new PercentileTDigestAccumulator(compressionFactor);
361+
aggregationResultHolder.setValue(accumulator);
362+
}
363+
return accumulator;
364+
}
365+
366+
private static PercentileTDigestAccumulator getSerializedAccumulator(GroupByResultHolder groupByResultHolder,
367+
int groupKey, byte[] bytes) {
368+
PercentileTDigestAccumulator accumulator = groupByResultHolder.getResult(groupKey);
369+
if (accumulator == null) {
370+
accumulator = PercentileTDigestAccumulator.forSerializedTDigest(bytes);
371+
groupByResultHolder.setValueForKey(groupKey, accumulator);
372+
}
373+
return accumulator;
374+
}
375+
376+
private static PercentileTDigestAccumulator getSerializedAccumulator(GroupByResultHolder groupByResultHolder,
377+
int groupKey, PercentileTDigestAccumulator.SerializedTDigestInput input) {
378+
PercentileTDigestAccumulator accumulator = groupByResultHolder.getResult(groupKey);
379+
if (accumulator == null) {
380+
accumulator = PercentileTDigestAccumulator.forSerializedTDigest(input);
381+
groupByResultHolder.setValueForKey(groupKey, accumulator);
382+
}
383+
return accumulator;
384+
}
385+
356386
/**
357387
* Returns the TDigest for the given group key if exists, or creates a new one with default compression.
358388
*

pinot-core/src/test/java/org/apache/pinot/core/common/ObjectSerDeUtilsTest.java

Lines changed: 60 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -20,6 +20,7 @@
2020

2121
import com.clearspring.analytics.stream.cardinality.HyperLogLog;
2222
import com.dynatrace.hash4j.distinctcount.UltraLogLog;
23+
import com.tdunning.math.stats.Centroid;
2324
import com.tdunning.math.stats.TDigest;
2425
import it.unimi.dsi.fastutil.doubles.Double2LongOpenHashMap;
2526
import it.unimi.dsi.fastutil.doubles.DoubleArrayList;
@@ -33,7 +34,9 @@
3334
import it.unimi.dsi.fastutil.longs.LongArrayList;
3435
import it.unimi.dsi.fastutil.objects.ObjectArrayList;
3536
import it.unimi.dsi.fastutil.objects.ObjectLinkedOpenHashSet;
37+
import java.nio.ByteBuffer;
3638
import java.util.Arrays;
39+
import java.util.Base64;
3740
import java.util.HashMap;
3841
import java.util.List;
3942
import java.util.Map;
@@ -64,6 +67,7 @@
6467
import org.apache.pinot.segment.local.customobject.TupleIntSketchAccumulator;
6568
import org.apache.pinot.segment.local.customobject.ValueLongPair;
6669
import org.apache.pinot.segment.local.utils.UltraLogLogUtils;
70+
import org.testng.annotations.DataProvider;
6771
import org.testng.annotations.Test;
6872

6973
import static org.testng.Assert.assertEquals;
@@ -77,6 +81,23 @@ public class ObjectSerDeUtilsTest {
7781
private static final String ERROR_MESSAGE = "Random seed: " + RANDOM_SEED;
7882

7983
private static final int NUM_ITERATIONS = 100;
84+
private static final String TDIGEST_3_2_VERBOSE_COMPRESSION_20 =
85+
"AAAAAT+5mZmZmZmaP+zMzMzMzM1ANAAAAAAAAAAAABtAKgAAAAAAAD+5mZmZmZmaQBgAAAAAAAA/uZmZmZmZmkAcAAAAAAAAP7mZ"
86+
+ "mZmZmZpAHAAAAAAAAD+5mZmZmZmaQBwAAAAAAAA/uZmZmZmZmkAmAAAAAAAAP7mZmZmZmZpAKgAAAAAAAD+5mZmZmZmaQCYAAAAA"
87+
+ "AAA/uZmZmZmZmkAmAAAAAAAAP7mZmZmZmZpAJAAAAAAAAD+5mZmZmZmaQC4AAAAAAAA/uZmZmZmZmkAmAAAAAAAAP7mZmZmZmZpA"
88+
+ "LAAAAAAAAD+5mZmZmZmaQCIAAAAAAAA/uZmZmZmZmkAyAAAAAAAAP8VVVVVVVVZAMQAAAAAAAD/gAAAAAAAAQCwAAAAAAAA/4AAA"
89+
+ "AAAAAEAsAAAAAAAAP+AAAAAAAABAKAAAAAAAAD/gAAAAAAAAQCQAAAAAAAA/564UeuFHrkAUAAAAAAAAP+zMzMzMzM1AFAAAAAAA"
90+
+ "AD/szMzMzMzNQBwAAAAAAAA/7MzMzMzMzUAQAAAAAAAAP+zMzMzMzM1ACAAAAAAAAD/szMzMzMzNP/AAAAAAAAA/7MzMzMzMzT/w"
91+
+ "AAAAAAAAP+zMzMzMzM0=";
92+
private static final String TDIGEST_3_2_SMALL_COMPRESSION_1000 =
93+
"AAAAAj+INkk/eVmAQJLlmhJu+WhEegAAB9oTiABAP4AAADxBsko/gAAAPRQvNz+AAAA9fAMgP4AAAD20Elw/gAAAPexqgT+AAAA+"
94+
+ "E5aAP4AAAD4yP88/gAAAPlJFjz+AAAA+c73OP4AAAD6LYDU/gAAAPp2zoz+AAAA+sOdBP4AAAD7FClg/gAAAPtotkD+AAAA+8GMX"
95+
+ "P4AAAD8D32Y/gAAAPxArOT+AAAA/HSDzP4AAAD8qzbU/gAAAPzk/8j+AAAA/SIegP4AAAD9YtmU/gAAAP2nf1D+AAAA/fBmxP4AA"
96+
+ "AD+Hvh4/gAAAP5IRRj+AAAA/nRWBP4AAAD+o294/gAAAP7V3mj+AAAA/wv55P4AAAD/RiSw/gAAAP+Ez2T+AAAA/8h6sP4AAAEAC"
97+
+ "N00/gAAAQAwnIz+AAABAFveKP4AAAEAixUk/gAAAQC+yDj+AAABAPeWGP4AAAEBNjrY/gAAAQF7lsz+AAABAci3iP4AAAECD3G8/"
98+
+ "gAAAQI/1LD+AAABAnZ6aP4AAAECtJUM/gAAAQL7pkj+AAABA02aDP4AAAEDrOzE/gAAAQQOcVz+AAABBFDtQP4AAAEEoOWA/gAAA"
99+
+ "QUChLz+AAABBXvGHP4AAAEGCsOA/gAAAQZutuD+AAABBvS/yP4AAAEHr7EA/gAAAQhhGxT+AAABCTlsKP4AAAEKWasI/gAAAQvfH"
100+
+ "yz+AAABDgpJvP4AAAESXLNE=";
80101

81102
@Test
82103
public void testString() {
@@ -332,6 +353,45 @@ public void testTDigest() {
332353
}
333354
}
334355

356+
@DataProvider(name = "tdigest32Fixtures")
357+
public static Object[][] tdigest32Fixtures() {
358+
return new Object[][]{
359+
{TDIGEST_3_2_VERBOSE_COMPRESSION_20, 1, 20.0, 256L, 0.1, 0.9},
360+
{TDIGEST_3_2_SMALL_COMPRESSION_1000, 2, 1_000.0, 64L, 0.011822292565892623, 1209.4004609431959}
361+
};
362+
}
363+
364+
@Test(dataProvider = "tdigest32Fixtures")
365+
public void testTDigest32Fixtures(String base64Bytes, int expectedEncoding, double expectedCompression,
366+
long expectedSize, double expectedMin, double expectedMax) {
367+
byte[] bytes = Base64.getDecoder().decode(base64Bytes);
368+
assertEquals(ByteBuffer.wrap(bytes).getInt(), expectedEncoding);
369+
TDigest digest = ObjectSerDeUtils.deserialize(bytes, ObjectSerDeUtils.ObjectType.TDigest);
370+
assertEquals(digest.compression(), expectedCompression);
371+
assertEquals(digest.size(), expectedSize);
372+
assertEquals(digest.getMin(), expectedMin);
373+
assertEquals(digest.getMax(), expectedMax);
374+
375+
long centroidWeight = 0L;
376+
double previousMean = Double.NEGATIVE_INFINITY;
377+
for (Centroid centroid : digest.centroids()) {
378+
assertTrue(Double.isFinite(centroid.mean()));
379+
assertTrue(centroid.count() > 0);
380+
assertTrue(centroid.mean() + 1e-12 >= previousMean);
381+
centroidWeight += centroid.count();
382+
previousMean = centroid.mean();
383+
}
384+
assertEquals(centroidWeight, expectedSize);
385+
386+
double previousQuantile = Double.NEGATIVE_INFINITY;
387+
for (double quantile : new double[]{0.0, 0.5, 0.75, 0.95, 0.99, 1.0}) {
388+
double value = digest.quantile(quantile);
389+
assertTrue(Double.isFinite(value));
390+
assertTrue(value >= previousQuantile);
391+
previousQuantile = value;
392+
}
393+
}
394+
335395
@Test
336396
public void testInt2LongMap() {
337397
for (int i = 0; i < NUM_ITERATIONS; i++) {

0 commit comments

Comments
 (0)