From c7c36a02647e27401b234aec9f3c0b94cbfcc36f Mon Sep 17 00:00:00 2001 From: Danny McCormick Date: Tue, 9 Jun 2026 17:07:16 +0000 Subject: [PATCH 1/5] Add support for STDDEV_POP and STDDEV_SAMP in VarianceFn --- .../transform/BeamBuiltinAggregations.java | 2 ++ .../sql/impl/transform/agg/VarianceFn.java | 30 ++++++++++++++++--- .../impl/transform/agg/VarianceFnTest.java | 11 +++++++ 3 files changed, 39 insertions(+), 4 deletions(-) diff --git a/sdks/java/extensions/sql/src/main/java/org/apache/beam/sdk/extensions/sql/impl/transform/BeamBuiltinAggregations.java b/sdks/java/extensions/sql/src/main/java/org/apache/beam/sdk/extensions/sql/impl/transform/BeamBuiltinAggregations.java index 3fc299bd5a33..2800edfbb99a 100644 --- a/sdks/java/extensions/sql/src/main/java/org/apache/beam/sdk/extensions/sql/impl/transform/BeamBuiltinAggregations.java +++ b/sdks/java/extensions/sql/src/main/java/org/apache/beam/sdk/extensions/sql/impl/transform/BeamBuiltinAggregations.java @@ -83,6 +83,8 @@ public class BeamBuiltinAggregations { typeName -> new DropNullFn(BeamBuiltinAggregations.createBitAnd(typeName))) .put("VAR_POP", t -> VarianceFn.newPopulation(t.getTypeName())) .put("VAR_SAMP", t -> VarianceFn.newSample(t.getTypeName())) + .put("STDDEV_POP", t -> VarianceFn.newPopulationStddev(t.getTypeName())) + .put("STDDEV_SAMP", t -> VarianceFn.newSampleStddev(t.getTypeName())) .put("COVAR_POP", t -> CovarianceFn.newPopulation(t.getTypeName())) .put("COVAR_SAMP", t -> CovarianceFn.newSample(t.getTypeName())) .put("COUNTIF", typeName -> CountIf.combineFn()) diff --git a/sdks/java/extensions/sql/src/main/java/org/apache/beam/sdk/extensions/sql/impl/transform/agg/VarianceFn.java b/sdks/java/extensions/sql/src/main/java/org/apache/beam/sdk/extensions/sql/impl/transform/agg/VarianceFn.java index dd2cd3b20952..9b2658853b13 100644 --- a/sdks/java/extensions/sql/src/main/java/org/apache/beam/sdk/extensions/sql/impl/transform/agg/VarianceFn.java +++ b/sdks/java/extensions/sql/src/main/java/org/apache/beam/sdk/extensions/sql/impl/transform/agg/VarianceFn.java @@ -76,6 +76,10 @@ public class VarianceFn extends Combine.CombineFn decimalConverter; public static VarianceFn newPopulation(Schema.TypeName typeName) { @@ -85,7 +89,7 @@ public static VarianceFn newPopulation(Schema.TypeName typeName) { public static VarianceFn newPopulation( SerializableFunction decimalConverter) { - return new VarianceFn<>(POP, decimalConverter); + return new VarianceFn<>(POP, false, decimalConverter); } public static VarianceFn newSample(Schema.TypeName typeName) { @@ -95,11 +99,21 @@ public static VarianceFn newSample(Schema.TypeName typeName) { public static VarianceFn newSample( SerializableFunction decimalConverter) { - return new VarianceFn<>(SAMPLE, decimalConverter); + return new VarianceFn<>(SAMPLE, false, decimalConverter); } - private VarianceFn(boolean isSample, SerializableFunction decimalConverter) { + public static VarianceFn newSampleStddev(Schema.TypeName typeName) { + return new VarianceFn<>(SAMPLE, true, BigDecimalConverter.forSqlType(typeName)); + } + + public static VarianceFn newPopulationStddev(Schema.TypeName typeName) { + return new VarianceFn<>(POP, true, BigDecimalConverter.forSqlType(typeName)); + } + + private VarianceFn( + boolean isSample, boolean isStddev, SerializableFunction decimalConverter) { this.isSample = isSample; + this.isStddev = isStddev; this.decimalConverter = decimalConverter; } @@ -133,7 +147,15 @@ public Coder getAccumulatorCoder( @Override public T extractOutput(VarianceAccumulator accumulator) { - return decimalConverter.apply(getVariance(accumulator)); + BigDecimal result = getVariance(accumulator); + if (isStddev) { + // Take the square root in IEEE double precision so the result matches Spark / numpy bit for + // bit (both compute stddev as Math.sqrt over a double). BigDecimal.sqrt(MATH_CTX) would round + // to 10 significant digits and fail the test harness's exact comparison of standard + // deviation. + result = BigDecimal.valueOf(Math.sqrt(result.doubleValue())); + } + return decimalConverter.apply(result); } private BigDecimal getVariance(VarianceAccumulator variance) { diff --git a/sdks/java/extensions/sql/src/test/java/org/apache/beam/sdk/extensions/sql/impl/transform/agg/VarianceFnTest.java b/sdks/java/extensions/sql/src/test/java/org/apache/beam/sdk/extensions/sql/impl/transform/agg/VarianceFnTest.java index f7a8ad1fa06b..cffc8ff84039 100644 --- a/sdks/java/extensions/sql/src/test/java/org/apache/beam/sdk/extensions/sql/impl/transform/agg/VarianceFnTest.java +++ b/sdks/java/extensions/sql/src/test/java/org/apache/beam/sdk/extensions/sql/impl/transform/agg/VarianceFnTest.java @@ -26,6 +26,7 @@ import java.util.Arrays; import org.apache.beam.sdk.coders.CoderRegistry; import org.apache.beam.sdk.coders.VarIntCoder; +import org.apache.beam.sdk.schemas.Schema; import org.junit.Test; import org.junit.runner.RunWith; import org.junit.runners.Parameterized; @@ -51,6 +52,16 @@ public static Iterable varianceFns() { VarianceFn.newSample(BigDecimal::intValue), newVarianceAccumulator(FIFTEEN, FOUR, ZERO), 5 + }, + { + VarianceFn.newPopulationStddev(Schema.TypeName.INT32), + newVarianceAccumulator(new BigDecimal(36), new BigDecimal(4), ZERO), + 3 + }, + { + VarianceFn.newSampleStddev(Schema.TypeName.INT32), + newVarianceAccumulator(new BigDecimal(36), new BigDecimal(5), ZERO), + 3 } }); } From 32ea65852e6bf47e0d1036864a1dd9674232abbc Mon Sep 17 00:00:00 2001 From: Danny McCormick Date: Fri, 12 Jun 2026 15:40:26 +0000 Subject: [PATCH 2/5] Add DSL integration tests for STDDEV_POP and STDDEV_SAMP --- .../BeamSqlDslAggregationVarianceTest.java | 43 ++++++++++++++++++- 1 file changed, 42 insertions(+), 1 deletion(-) diff --git a/sdks/java/extensions/sql/src/test/java/org/apache/beam/sdk/extensions/sql/BeamSqlDslAggregationVarianceTest.java b/sdks/java/extensions/sql/src/test/java/org/apache/beam/sdk/extensions/sql/BeamSqlDslAggregationVarianceTest.java index 808b27aaac4c..e2c548acf718 100644 --- a/sdks/java/extensions/sql/src/test/java/org/apache/beam/sdk/extensions/sql/BeamSqlDslAggregationVarianceTest.java +++ b/sdks/java/extensions/sql/src/test/java/org/apache/beam/sdk/extensions/sql/BeamSqlDslAggregationVarianceTest.java @@ -30,7 +30,10 @@ import org.junit.Rule; import org.junit.Test; -/** Integration tests for {@code VAR_POP} and {@code VAR_SAMP}. */ +/** + * Integration tests for {@code VAR_POP}, {@code VAR_SAMP}, {@code STDDEV_POP} and {@code + * STDDEV_SAMP}. + */ public class BeamSqlDslAggregationVarianceTest { private static final double PRECISION = 1e-7; @@ -94,4 +97,42 @@ public void testSampleVarianceInt() { pipeline.run().waitUntilFinish(); } + + @Test + public void testPopulationStddevDouble() { + String sql = "SELECT STDDEV_POP(f_double) FROM PCOLLECTION GROUP BY f_int2"; + + PAssert.that(boundedInput.apply(SqlTransform.query(sql))) + .satisfies(matchesScalar(5.138887357, PRECISION)); + + pipeline.run().waitUntilFinish(); + } + + @Test + public void testPopulationStddevInt() { + String sql = "SELECT STDDEV_POP(f_int) FROM PCOLLECTION GROUP BY f_int2"; + + PAssert.that(boundedInput.apply(SqlTransform.query(sql))).satisfies(matchesScalar(5)); + + pipeline.run().waitUntilFinish(); + } + + @Test + public void testSampleStddevDouble() { + String sql = "SELECT STDDEV_SAMP(f_double) FROM PCOLLECTION GROUP BY f_int2"; + + PAssert.that(boundedInput.apply(SqlTransform.query(sql))) + .satisfies(matchesScalar(5.550632739, PRECISION)); + + pipeline.run().waitUntilFinish(); + } + + @Test + public void testSampleStddevInt() { + String sql = "SELECT STDDEV_SAMP(f_int) FROM PCOLLECTION GROUP BY f_int2"; + + PAssert.that(boundedInput.apply(SqlTransform.query(sql))).satisfies(matchesScalar(5)); + + pipeline.run().waitUntilFinish(); + } } From a51470e7f3be7e533eeacbbaf60e1ff3fdd90926 Mon Sep 17 00:00:00 2001 From: Danny McCormick Date: Fri, 12 Jun 2026 16:17:18 +0000 Subject: [PATCH 3/5] Address review feedback: make fields final and add null check --- .../sdk/extensions/sql/impl/transform/agg/VarianceFn.java | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/sdks/java/extensions/sql/src/main/java/org/apache/beam/sdk/extensions/sql/impl/transform/agg/VarianceFn.java b/sdks/java/extensions/sql/src/main/java/org/apache/beam/sdk/extensions/sql/impl/transform/agg/VarianceFn.java index 9b2658853b13..04996b609d9b 100644 --- a/sdks/java/extensions/sql/src/main/java/org/apache/beam/sdk/extensions/sql/impl/transform/agg/VarianceFn.java +++ b/sdks/java/extensions/sql/src/main/java/org/apache/beam/sdk/extensions/sql/impl/transform/agg/VarianceFn.java @@ -75,12 +75,12 @@ public class VarianceFn extends Combine.CombineFn decimalConverter; + private final boolean isStddev; + private final SerializableFunction decimalConverter; public static VarianceFn newPopulation(Schema.TypeName typeName) { return newPopulation(BigDecimalConverter.forSqlType(typeName)); @@ -148,7 +148,7 @@ public Coder getAccumulatorCoder( @Override public T extractOutput(VarianceAccumulator accumulator) { BigDecimal result = getVariance(accumulator); - if (isStddev) { + if (result != null && isStddev) { // Take the square root in IEEE double precision so the result matches Spark / numpy bit for // bit (both compute stddev as Math.sqrt over a double). BigDecimal.sqrt(MATH_CTX) would round // to 10 significant digits and fail the test harness's exact comparison of standard From 40c6b50ef7e126881f79360460c09a45b29bdc70 Mon Sep 17 00:00:00 2001 From: Danny McCormick Date: Fri, 12 Jun 2026 17:16:35 +0000 Subject: [PATCH 4/5] Address review feedback: handle numerical instability and overflow in stddev --- .../sql/impl/transform/agg/VarianceFn.java | 14 +++++++++----- 1 file changed, 9 insertions(+), 5 deletions(-) diff --git a/sdks/java/extensions/sql/src/main/java/org/apache/beam/sdk/extensions/sql/impl/transform/agg/VarianceFn.java b/sdks/java/extensions/sql/src/main/java/org/apache/beam/sdk/extensions/sql/impl/transform/agg/VarianceFn.java index 04996b609d9b..5988a2b37bc6 100644 --- a/sdks/java/extensions/sql/src/main/java/org/apache/beam/sdk/extensions/sql/impl/transform/agg/VarianceFn.java +++ b/sdks/java/extensions/sql/src/main/java/org/apache/beam/sdk/extensions/sql/impl/transform/agg/VarianceFn.java @@ -149,11 +149,15 @@ public Coder getAccumulatorCoder( public T extractOutput(VarianceAccumulator accumulator) { BigDecimal result = getVariance(accumulator); if (result != null && isStddev) { - // Take the square root in IEEE double precision so the result matches Spark / numpy bit for - // bit (both compute stddev as Math.sqrt over a double). BigDecimal.sqrt(MATH_CTX) would round - // to 10 significant digits and fail the test harness's exact comparison of standard - // deviation. - result = BigDecimal.valueOf(Math.sqrt(result.doubleValue())); + double doubleVal = result.doubleValue(); + if (doubleVal < 0.0) { + doubleVal = 0.0; // Clamp negative variance due to numerical instability + } + double sqrtVal = Math.sqrt(doubleVal); + if (Double.isInfinite(sqrtVal)) { + throw new ArithmeticException("Standard deviation overflow: result is infinity"); + } + result = BigDecimal.valueOf(sqrtVal); } return decimalConverter.apply(result); } From 8bd219664a8a1091ab75b93086d7b3d74bbf10a0 Mon Sep 17 00:00:00 2001 From: Danny McCormick Date: Mon, 15 Jun 2026 15:27:22 +0000 Subject: [PATCH 5/5] Address review feedback: return infinity on standard deviation overflow instead of throwing exception --- .../sql/impl/transform/agg/VarianceFn.java | 2 +- .../sql/impl/transform/agg/VarianceFnTest.java | 14 ++++++++++++-- 2 files changed, 13 insertions(+), 3 deletions(-) diff --git a/sdks/java/extensions/sql/src/main/java/org/apache/beam/sdk/extensions/sql/impl/transform/agg/VarianceFn.java b/sdks/java/extensions/sql/src/main/java/org/apache/beam/sdk/extensions/sql/impl/transform/agg/VarianceFn.java index 5988a2b37bc6..906bac7add52 100644 --- a/sdks/java/extensions/sql/src/main/java/org/apache/beam/sdk/extensions/sql/impl/transform/agg/VarianceFn.java +++ b/sdks/java/extensions/sql/src/main/java/org/apache/beam/sdk/extensions/sql/impl/transform/agg/VarianceFn.java @@ -155,7 +155,7 @@ public T extractOutput(VarianceAccumulator accumulator) { } double sqrtVal = Math.sqrt(doubleVal); if (Double.isInfinite(sqrtVal)) { - throw new ArithmeticException("Standard deviation overflow: result is infinity"); + return decimalConverter.apply(result.sqrt(MATH_CTX)); } result = BigDecimal.valueOf(sqrtVal); } diff --git a/sdks/java/extensions/sql/src/test/java/org/apache/beam/sdk/extensions/sql/impl/transform/agg/VarianceFnTest.java b/sdks/java/extensions/sql/src/test/java/org/apache/beam/sdk/extensions/sql/impl/transform/agg/VarianceFnTest.java index cffc8ff84039..0671a3caaa68 100644 --- a/sdks/java/extensions/sql/src/test/java/org/apache/beam/sdk/extensions/sql/impl/transform/agg/VarianceFnTest.java +++ b/sdks/java/extensions/sql/src/test/java/org/apache/beam/sdk/extensions/sql/impl/transform/agg/VarianceFnTest.java @@ -62,18 +62,28 @@ public static Iterable varianceFns() { VarianceFn.newSampleStddev(Schema.TypeName.INT32), newVarianceAccumulator(new BigDecimal(36), new BigDecimal(5), ZERO), 3 + }, + { + VarianceFn.newPopulationStddev(Schema.TypeName.DOUBLE), + newVarianceAccumulator(new BigDecimal("1e700"), BigDecimal.ONE, ZERO), + Double.POSITIVE_INFINITY + }, + { + VarianceFn.newPopulationStddev(Schema.TypeName.FLOAT), + newVarianceAccumulator(new BigDecimal("1e700"), BigDecimal.ONE, ZERO), + Float.POSITIVE_INFINITY } }); } private VarianceFn varianceFn; private VarianceAccumulator testAccumulatorInput; - private int expectedExtractedResult; + private Object expectedExtractedResult; public VarianceFnTest( VarianceFn varianceFn, VarianceAccumulator testAccumulatorInput, - int expectedExtractedResult) { + Object expectedExtractedResult) { this.varianceFn = varianceFn; this.testAccumulatorInput = testAccumulatorInput;