Skip to content

Commit bde545f

Browse files
authored
Add support for STDDEV_POP and STDDEV_SAMP in VarianceFn (#38871)
* Add support for STDDEV_POP and STDDEV_SAMP in VarianceFn * Add DSL integration tests for STDDEV_POP and STDDEV_SAMP * Address review feedback: make fields final and add null check * Address review feedback: handle numerical instability and overflow in stddev * Address review feedback: return infinity on standard deviation overflow instead of throwing exception
1 parent 98e29d9 commit bde545f

4 files changed

Lines changed: 99 additions & 9 deletions

File tree

sdks/java/extensions/sql/src/main/java/org/apache/beam/sdk/extensions/sql/impl/transform/BeamBuiltinAggregations.java

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -83,6 +83,8 @@ public class BeamBuiltinAggregations {
8383
typeName -> new DropNullFn(BeamBuiltinAggregations.createBitAnd(typeName)))
8484
.put("VAR_POP", t -> VarianceFn.newPopulation(t.getTypeName()))
8585
.put("VAR_SAMP", t -> VarianceFn.newSample(t.getTypeName()))
86+
.put("STDDEV_POP", t -> VarianceFn.newPopulationStddev(t.getTypeName()))
87+
.put("STDDEV_SAMP", t -> VarianceFn.newSampleStddev(t.getTypeName()))
8688
.put("COVAR_POP", t -> CovarianceFn.newPopulation(t.getTypeName()))
8789
.put("COVAR_SAMP", t -> CovarianceFn.newSample(t.getTypeName()))
8890
.put("COUNTIF", typeName -> CountIf.combineFn())

sdks/java/extensions/sql/src/main/java/org/apache/beam/sdk/extensions/sql/impl/transform/agg/VarianceFn.java

Lines changed: 32 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -75,8 +75,12 @@ public class VarianceFn<T extends Number> extends Combine.CombineFn<T, VarianceA
7575
private static final boolean SAMPLE = true;
7676
private static final boolean POP = false;
7777

78-
private boolean isSample; // flag to determine return value should be Variance Pop or Sample
79-
private SerializableFunction<BigDecimal, T> decimalConverter;
78+
private final boolean isSample; // flag to determine return value should be Variance Pop or Sample
79+
// When true, extractOutput returns the square root of the variance (i.e. standard deviation).
80+
// Beam's enumerable bridge cannot translate a SQRT call layered on top of a window VAR_SAMP, so
81+
// STDDEV_SAMP / STDDEV_POP are computed end-to-end inside this combiner instead.
82+
private final boolean isStddev;
83+
private final SerializableFunction<BigDecimal, T> decimalConverter;
8084

8185
public static VarianceFn newPopulation(Schema.TypeName typeName) {
8286
return newPopulation(BigDecimalConverter.forSqlType(typeName));
@@ -85,7 +89,7 @@ public static VarianceFn newPopulation(Schema.TypeName typeName) {
8589
public static <V extends Number> VarianceFn newPopulation(
8690
SerializableFunction<BigDecimal, V> decimalConverter) {
8791

88-
return new VarianceFn<>(POP, decimalConverter);
92+
return new VarianceFn<>(POP, false, decimalConverter);
8993
}
9094

9195
public static VarianceFn newSample(Schema.TypeName typeName) {
@@ -95,11 +99,21 @@ public static VarianceFn newSample(Schema.TypeName typeName) {
9599
public static <V extends Number> VarianceFn newSample(
96100
SerializableFunction<BigDecimal, V> decimalConverter) {
97101

98-
return new VarianceFn<>(SAMPLE, decimalConverter);
102+
return new VarianceFn<>(SAMPLE, false, decimalConverter);
99103
}
100104

101-
private VarianceFn(boolean isSample, SerializableFunction<BigDecimal, T> decimalConverter) {
105+
public static VarianceFn newSampleStddev(Schema.TypeName typeName) {
106+
return new VarianceFn<>(SAMPLE, true, BigDecimalConverter.forSqlType(typeName));
107+
}
108+
109+
public static VarianceFn newPopulationStddev(Schema.TypeName typeName) {
110+
return new VarianceFn<>(POP, true, BigDecimalConverter.forSqlType(typeName));
111+
}
112+
113+
private VarianceFn(
114+
boolean isSample, boolean isStddev, SerializableFunction<BigDecimal, T> decimalConverter) {
102115
this.isSample = isSample;
116+
this.isStddev = isStddev;
103117
this.decimalConverter = decimalConverter;
104118
}
105119

@@ -133,7 +147,19 @@ public Coder<VarianceAccumulator> getAccumulatorCoder(
133147

134148
@Override
135149
public T extractOutput(VarianceAccumulator accumulator) {
136-
return decimalConverter.apply(getVariance(accumulator));
150+
BigDecimal result = getVariance(accumulator);
151+
if (result != null && isStddev) {
152+
double doubleVal = result.doubleValue();
153+
if (doubleVal < 0.0) {
154+
doubleVal = 0.0; // Clamp negative variance due to numerical instability
155+
}
156+
double sqrtVal = Math.sqrt(doubleVal);
157+
if (Double.isInfinite(sqrtVal)) {
158+
return decimalConverter.apply(result.sqrt(MATH_CTX));
159+
}
160+
result = BigDecimal.valueOf(sqrtVal);
161+
}
162+
return decimalConverter.apply(result);
137163
}
138164

139165
private BigDecimal getVariance(VarianceAccumulator variance) {

sdks/java/extensions/sql/src/test/java/org/apache/beam/sdk/extensions/sql/BeamSqlDslAggregationVarianceTest.java

Lines changed: 42 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -30,7 +30,10 @@
3030
import org.junit.Rule;
3131
import org.junit.Test;
3232

33-
/** Integration tests for {@code VAR_POP} and {@code VAR_SAMP}. */
33+
/**
34+
* Integration tests for {@code VAR_POP}, {@code VAR_SAMP}, {@code STDDEV_POP} and {@code
35+
* STDDEV_SAMP}.
36+
*/
3437
public class BeamSqlDslAggregationVarianceTest {
3538

3639
private static final double PRECISION = 1e-7;
@@ -94,4 +97,42 @@ public void testSampleVarianceInt() {
9497

9598
pipeline.run().waitUntilFinish();
9699
}
100+
101+
@Test
102+
public void testPopulationStddevDouble() {
103+
String sql = "SELECT STDDEV_POP(f_double) FROM PCOLLECTION GROUP BY f_int2";
104+
105+
PAssert.that(boundedInput.apply(SqlTransform.query(sql)))
106+
.satisfies(matchesScalar(5.138887357, PRECISION));
107+
108+
pipeline.run().waitUntilFinish();
109+
}
110+
111+
@Test
112+
public void testPopulationStddevInt() {
113+
String sql = "SELECT STDDEV_POP(f_int) FROM PCOLLECTION GROUP BY f_int2";
114+
115+
PAssert.that(boundedInput.apply(SqlTransform.query(sql))).satisfies(matchesScalar(5));
116+
117+
pipeline.run().waitUntilFinish();
118+
}
119+
120+
@Test
121+
public void testSampleStddevDouble() {
122+
String sql = "SELECT STDDEV_SAMP(f_double) FROM PCOLLECTION GROUP BY f_int2";
123+
124+
PAssert.that(boundedInput.apply(SqlTransform.query(sql)))
125+
.satisfies(matchesScalar(5.550632739, PRECISION));
126+
127+
pipeline.run().waitUntilFinish();
128+
}
129+
130+
@Test
131+
public void testSampleStddevInt() {
132+
String sql = "SELECT STDDEV_SAMP(f_int) FROM PCOLLECTION GROUP BY f_int2";
133+
134+
PAssert.that(boundedInput.apply(SqlTransform.query(sql))).satisfies(matchesScalar(5));
135+
136+
pipeline.run().waitUntilFinish();
137+
}
97138
}

sdks/java/extensions/sql/src/test/java/org/apache/beam/sdk/extensions/sql/impl/transform/agg/VarianceFnTest.java

Lines changed: 23 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -26,6 +26,7 @@
2626
import java.util.Arrays;
2727
import org.apache.beam.sdk.coders.CoderRegistry;
2828
import org.apache.beam.sdk.coders.VarIntCoder;
29+
import org.apache.beam.sdk.schemas.Schema;
2930
import org.junit.Test;
3031
import org.junit.runner.RunWith;
3132
import org.junit.runners.Parameterized;
@@ -51,18 +52,38 @@ public static Iterable<Object[]> varianceFns() {
5152
VarianceFn.newSample(BigDecimal::intValue),
5253
newVarianceAccumulator(FIFTEEN, FOUR, ZERO),
5354
5
55+
},
56+
{
57+
VarianceFn.newPopulationStddev(Schema.TypeName.INT32),
58+
newVarianceAccumulator(new BigDecimal(36), new BigDecimal(4), ZERO),
59+
3
60+
},
61+
{
62+
VarianceFn.newSampleStddev(Schema.TypeName.INT32),
63+
newVarianceAccumulator(new BigDecimal(36), new BigDecimal(5), ZERO),
64+
3
65+
},
66+
{
67+
VarianceFn.newPopulationStddev(Schema.TypeName.DOUBLE),
68+
newVarianceAccumulator(new BigDecimal("1e700"), BigDecimal.ONE, ZERO),
69+
Double.POSITIVE_INFINITY
70+
},
71+
{
72+
VarianceFn.newPopulationStddev(Schema.TypeName.FLOAT),
73+
newVarianceAccumulator(new BigDecimal("1e700"), BigDecimal.ONE, ZERO),
74+
Float.POSITIVE_INFINITY
5475
}
5576
});
5677
}
5778

5879
private VarianceFn varianceFn;
5980
private VarianceAccumulator testAccumulatorInput;
60-
private int expectedExtractedResult;
81+
private Object expectedExtractedResult;
6182

6283
public VarianceFnTest(
6384
VarianceFn varianceFn,
6485
VarianceAccumulator testAccumulatorInput,
65-
int expectedExtractedResult) {
86+
Object expectedExtractedResult) {
6687

6788
this.varianceFn = varianceFn;
6889
this.testAccumulatorInput = testAccumulatorInput;

0 commit comments

Comments
 (0)