This is an automated email from the ASF dual-hosted git repository. jt2594838 pushed a commit to branch remove_swtich_type in repository https://gitbox.apache.org/repos/asf/iotdb.git
commit b239e8ae5772e68f7c0412a5f164204974785de5 Author: Tian Jiang <[email protected]> AuthorDate: Mon Aug 24 10:45:59 2026 +0800 multiple refactors --- .../org/apache/iotdb/calc/i18n/CalcMessages.java | 4 + .../org/apache/iotdb/calc/i18n/CalcMessages.java | 4 + .../execution/aggregation/VarianceAccumulator.java | 116 +++-------- .../aggregation/TableVarianceAccumulator.java | 223 +++------------------ .../grouped/GroupedVarianceAccumulator.java | 140 ++----------- .../aggregation/VarianceAccumulatorTest.java | 126 ++++++++++++ 6 files changed, 215 insertions(+), 398 deletions(-) diff --git a/iotdb-core/calc-commons/src/main/i18n/en/org/apache/iotdb/calc/i18n/CalcMessages.java b/iotdb-core/calc-commons/src/main/i18n/en/org/apache/iotdb/calc/i18n/CalcMessages.java index c5c339d6df5..688797980fc 100644 --- a/iotdb-core/calc-commons/src/main/i18n/en/org/apache/iotdb/calc/i18n/CalcMessages.java +++ b/iotdb-core/calc-commons/src/main/i18n/en/org/apache/iotdb/calc/i18n/CalcMessages.java @@ -135,6 +135,10 @@ public final class CalcMessages { public static final String UNSUPPORTED_DATA_TYPE = "Unsupported data type: "; public static final String UNSUPPORTED_DATA_TYPE_IN_CENTRAL_MOMENT_AGGREGATION = "Unsupported data type in CentralMoment Aggregation: %s"; + public static final String UNSUPPORTED_DATA_TYPE_IN_AGGREGATION_VARIANCE = + "Unsupported data type in aggregation variance : %s"; + public static final String UNSUPPORTED_DATA_TYPE_IN_VARIANCE_AGGREGATION = + "Unsupported data type in VARIANCE Aggregation: %s"; public static final String UNSUPPORTED_DEFAULT_VALUE_DATA_TYPE_IN_LAG = "Unsupported default value's data type in Lag: "; public static final String UNSUPPORTED_DATA_TYPE_LOWER = "unsupported data type: "; diff --git a/iotdb-core/calc-commons/src/main/i18n/zh/org/apache/iotdb/calc/i18n/CalcMessages.java b/iotdb-core/calc-commons/src/main/i18n/zh/org/apache/iotdb/calc/i18n/CalcMessages.java index 05a8e667fbf..84088a79202 100644 --- a/iotdb-core/calc-commons/src/main/i18n/zh/org/apache/iotdb/calc/i18n/CalcMessages.java +++ b/iotdb-core/calc-commons/src/main/i18n/zh/org/apache/iotdb/calc/i18n/CalcMessages.java @@ -128,6 +128,10 @@ public final class CalcMessages { public static final String UNSUPPORTED_DATA_TYPE = "不支持的数据类型:"; public static final String UNSUPPORTED_DATA_TYPE_IN_CENTRAL_MOMENT_AGGREGATION = "CentralMoment 聚合中不支持的数据类型:%s"; + public static final String UNSUPPORTED_DATA_TYPE_IN_AGGREGATION_VARIANCE = + "variance 聚合中不支持的数据类型:%s"; + public static final String UNSUPPORTED_DATA_TYPE_IN_VARIANCE_AGGREGATION = + "VARIANCE 聚合中不支持的数据类型:%s"; public static final String UNSUPPORTED_DEFAULT_VALUE_DATA_TYPE_IN_LAG = "Lag 中不支持的默认值数据类型:"; public static final String UNSUPPORTED_DATA_TYPE_LOWER = "不支持的数据类型:"; diff --git a/iotdb-core/calc-commons/src/main/java/org/apache/iotdb/calc/execution/aggregation/VarianceAccumulator.java b/iotdb-core/calc-commons/src/main/java/org/apache/iotdb/calc/execution/aggregation/VarianceAccumulator.java index 48972749c3a..6c83d70eca7 100644 --- a/iotdb-core/calc-commons/src/main/java/org/apache/iotdb/calc/execution/aggregation/VarianceAccumulator.java +++ b/iotdb-core/calc-commons/src/main/java/org/apache/iotdb/calc/execution/aggregation/VarianceAccumulator.java @@ -19,10 +19,14 @@ package org.apache.iotdb.calc.execution.aggregation; +import org.apache.iotdb.calc.i18n.CalcMessages; +import org.apache.iotdb.calc.utils.TypeServices; + import org.apache.tsfile.block.column.Column; import org.apache.tsfile.block.column.ColumnBuilder; import org.apache.tsfile.enums.TSDataType; import org.apache.tsfile.file.metadata.statistics.Statistics; +import org.apache.tsfile.read.common.type.Type; import org.apache.tsfile.utils.Binary; import org.apache.tsfile.utils.BitMap; import org.apache.tsfile.utils.BytesUtils; @@ -41,6 +45,7 @@ public class VarianceAccumulator implements Accumulator { } private final TSDataType seriesDataType; + private final TypeServices.ColumnToDoubleConverter doubleValueConverter; private final VarianceType varianceType; @@ -50,34 +55,33 @@ public class VarianceAccumulator implements Accumulator { public VarianceAccumulator(TSDataType seriesDataType, VarianceType varianceType) { this.seriesDataType = seriesDataType; + this.doubleValueConverter = + TypeServices.NUMERIC_COLUMN_TO_DOUBLE_CONVERTER_SERVICE + .call(Type.fromTsDataType(seriesDataType)) + .create( + () -> + new UnSupportedDataTypeException( + String.format( + CalcMessages.UNSUPPORTED_DATA_TYPE_IN_AGGREGATION_VARIANCE, + seriesDataType))); this.varianceType = varianceType; } @Override public void addInput(Column[] columns, BitMap bitMap) { - switch (seriesDataType) { - case INT32: - addIntInput(columns, bitMap); - return; - case INT64: - addLongInput(columns, bitMap); - return; - case FLOAT: - addFloatInput(columns, bitMap); - return; - case DOUBLE: - addDoubleInput(columns, bitMap); - return; - case TEXT: - case BLOB: - case OBJECT: - case BOOLEAN: - case DATE: - case STRING: - case TIMESTAMP: - default: - throw new UnSupportedDataTypeException( - String.format("Unsupported data type in aggregation variance : %s", seriesDataType)); + checkInputDataType(); + int size = columns[0].getPositionCount(); + for (int i = 0; i < size; i++) { + if (bitMap != null && !bitMap.isMarked(i)) { + continue; + } + if (!columns[1].isNull(i)) { + double value = doubleValueConverter.convert(columns[1], i); + count++; + double delta = value - mean; + mean += delta / count; + m2 += delta * (value - mean); + } } } @@ -216,67 +220,11 @@ public class VarianceAccumulator implements Accumulator { return TSDataType.DOUBLE; } - private void addIntInput(Column[] columns, BitMap bitmap) { - int size = columns[0].getPositionCount(); - for (int i = 0; i < size; i++) { - if (bitmap != null && !bitmap.isMarked(i)) { - continue; - } - if (!columns[1].isNull(i)) { - int value = columns[1].getInt(i); - count++; - double delta = value - mean; - mean += delta / count; - m2 += delta * (value - mean); - } - } - } - - private void addLongInput(Column[] columns, BitMap bitmap) { - int size = columns[0].getPositionCount(); - for (int i = 0; i < size; i++) { - if (bitmap != null && !bitmap.isMarked(i)) { - continue; - } - if (!columns[1].isNull(i)) { - long value = columns[1].getLong(i); - count++; - double delta = value - mean; - mean += delta / count; - m2 += delta * (value - mean); - } - } - } - - private void addFloatInput(Column[] columns, BitMap bitmap) { - int size = columns[0].getPositionCount(); - for (int i = 0; i < size; i++) { - if (bitmap != null && !bitmap.isMarked(i)) { - continue; - } - if (!columns[1].isNull(i)) { - float value = columns[1].getFloat(i); - count++; - double delta = value - mean; - mean += delta / count; - m2 += delta * (value - mean); - } - } - } - - private void addDoubleInput(Column[] columns, BitMap bitmap) { - int size = columns[0].getPositionCount(); - for (int i = 0; i < size; i++) { - if (bitmap != null && !bitmap.isMarked(i)) { - continue; - } - if (!columns[1].isNull(i)) { - double value = columns[1].getDouble(i); - count++; - double delta = value - mean; - mean += delta / count; - m2 += delta * (value - mean); - } + private void checkInputDataType() { + if (!seriesDataType.isNumeric()) { + throw new UnSupportedDataTypeException( + String.format( + CalcMessages.UNSUPPORTED_DATA_TYPE_IN_AGGREGATION_VARIANCE, seriesDataType)); } } } diff --git a/iotdb-core/calc-commons/src/main/java/org/apache/iotdb/calc/execution/operator/source/relational/aggregation/TableVarianceAccumulator.java b/iotdb-core/calc-commons/src/main/java/org/apache/iotdb/calc/execution/operator/source/relational/aggregation/TableVarianceAccumulator.java index ef186d8d3fb..c8efef831b2 100644 --- a/iotdb-core/calc-commons/src/main/java/org/apache/iotdb/calc/execution/operator/source/relational/aggregation/TableVarianceAccumulator.java +++ b/iotdb-core/calc-commons/src/main/java/org/apache/iotdb/calc/execution/operator/source/relational/aggregation/TableVarianceAccumulator.java @@ -20,6 +20,8 @@ package org.apache.iotdb.calc.execution.operator.source.relational.aggregation; import org.apache.iotdb.calc.execution.aggregation.VarianceAccumulator; +import org.apache.iotdb.calc.i18n.CalcMessages; +import org.apache.iotdb.calc.utils.TypeServices; import org.apache.tsfile.block.column.Column; import org.apache.tsfile.block.column.ColumnBuilder; @@ -28,6 +30,7 @@ import org.apache.tsfile.file.metadata.statistics.Statistics; import org.apache.tsfile.read.common.block.column.BinaryColumn; import org.apache.tsfile.read.common.block.column.BinaryColumnBuilder; import org.apache.tsfile.read.common.block.column.RunLengthEncodedColumn; +import org.apache.tsfile.read.common.type.Type; import org.apache.tsfile.utils.Binary; import org.apache.tsfile.utils.BytesUtils; import org.apache.tsfile.utils.RamUsageEstimator; @@ -40,6 +43,7 @@ public class TableVarianceAccumulator implements TableAccumulator { private static final long INSTANCE_SIZE = RamUsageEstimator.shallowSizeOfInstance(TableVarianceAccumulator.class); private final TSDataType seriesDataType; + private final TypeServices.ColumnToDoubleConverter doubleValueConverter; private final VarianceAccumulator.VarianceType varianceType; private long count; @@ -49,6 +53,10 @@ public class TableVarianceAccumulator implements TableAccumulator { public TableVarianceAccumulator( TSDataType seriesDataType, VarianceAccumulator.VarianceType varianceType) { this.seriesDataType = seriesDataType; + this.doubleValueConverter = + TypeServices.NUMERIC_COLUMN_TO_DOUBLE_CONVERTER_SERVICE + .call(Type.fromTsDataType(seriesDataType)) + .create(this::unsupportedDataTypeException); this.varianceType = varianceType; } @@ -64,57 +72,18 @@ public class TableVarianceAccumulator implements TableAccumulator { @Override public void addInput(Column[] arguments, AggregationMask mask) { - switch (seriesDataType) { - case INT32: - addIntInput(arguments[0], mask); - return; - case INT64: - addLongInput(arguments[0], mask); - return; - case FLOAT: - addFloatInput(arguments[0], mask); - return; - case DOUBLE: - addDoubleInput(arguments[0], mask); - return; - case TEXT: - case BLOB: - case OBJECT: - case BOOLEAN: - case DATE: - case STRING: - case TIMESTAMP: - default: - throw new UnSupportedDataTypeException( - String.format("Unsupported data type in VARIANCE Aggregation: %s", seriesDataType)); - } + checkInputDataType(); + updateStateByAdd(arguments[0], mask); } @Override public void removeInput(Column[] arguments) { - switch (seriesDataType) { - case INT32: - removeIntInput(arguments[0]); - return; - case INT64: - removeLongInput(arguments[0]); - return; - case FLOAT: - removeFloatInput(arguments[0]); - return; - case DOUBLE: - removeDoubleInput(arguments[0]); - return; - case TEXT: - case BLOB: - case OBJECT: - case BOOLEAN: - case DATE: - case STRING: - case TIMESTAMP: - default: - throw new UnSupportedDataTypeException( - String.format("Unsupported data type in VARIANCE Aggregation: %s", seriesDataType)); + checkInputDataType(); + Column column = arguments[0]; + for (int i = 0; i < column.getPositionCount(); i++) { + if (!column.isNull(i)) { + updateStateByRemove(doubleValueConverter.convert(column, i)); + } } } @@ -222,106 +191,7 @@ public class TableVarianceAccumulator implements TableAccumulator { return true; } - private void addIntInput(Column column, AggregationMask mask) { - int positionCount = mask.getSelectedPositionCount(); - - if (mask.isSelectAll()) { - for (int i = 0; i < positionCount; i++) { - if (column.isNull(i)) { - continue; - } - - int value = column.getInt(i); - count++; - double delta = value - mean; - mean += delta / count; - m2 += delta * (value - mean); - } - } else { - int[] selectedPositions = mask.getSelectedPositions(); - int position; - for (int i = 0; i < positionCount; i++) { - position = selectedPositions[i]; - if (column.isNull(position)) { - continue; - } - - int value = column.getInt(position); - count++; - double delta = value - mean; - mean += delta / count; - m2 += delta * (value - mean); - } - } - } - - private void addLongInput(Column column, AggregationMask mask) { - int positionCount = mask.getSelectedPositionCount(); - - if (mask.isSelectAll()) { - for (int i = 0; i < positionCount; i++) { - if (column.isNull(i)) { - continue; - } - - long value = column.getLong(i); - count++; - double delta = value - mean; - mean += delta / count; - m2 += delta * (value - mean); - } - } else { - int[] selectedPositions = mask.getSelectedPositions(); - int position; - for (int i = 0; i < positionCount; i++) { - position = selectedPositions[i]; - if (column.isNull(position)) { - continue; - } - - long value = column.getLong(position); - count++; - double delta = value - mean; - mean += delta / count; - m2 += delta * (value - mean); - } - } - } - - private void addFloatInput(Column column, AggregationMask mask) { - int positionCount = mask.getSelectedPositionCount(); - - if (mask.isSelectAll()) { - for (int i = 0; i < positionCount; i++) { - if (column.isNull(i)) { - continue; - } - - float value = column.getFloat(i); - count++; - double delta = value - mean; - mean += delta / count; - m2 += delta * (value - mean); - } - } else { - int[] selectedPositions = mask.getSelectedPositions(); - int position; - for (int i = 0; i < positionCount; i++) { - position = selectedPositions[i]; - if (column.isNull(position)) { - continue; - } - - float value = column.getFloat(position); - count++; - double delta = value - mean; - mean += delta / count; - m2 += delta * (value - mean); - } - } - } - - private void addDoubleInput(Column column, AggregationMask mask) { + private void updateStateByAdd(Column column, AggregationMask mask) { int positionCount = mask.getSelectedPositionCount(); if (mask.isSelectAll()) { @@ -330,7 +200,7 @@ public class TableVarianceAccumulator implements TableAccumulator { continue; } - double value = column.getDouble(i); + double value = doubleValueConverter.convert(column, i); count++; double delta = value - mean; mean += delta / count; @@ -345,7 +215,7 @@ public class TableVarianceAccumulator implements TableAccumulator { continue; } - double value = column.getDouble(position); + double value = doubleValueConverter.convert(column, position); count++; double delta = value - mean; mean += delta / count; @@ -354,50 +224,6 @@ public class TableVarianceAccumulator implements TableAccumulator { } } - private void removeIntInput(Column column) { - for (int i = 0; i < column.getPositionCount(); i++) { - if (column.isNull(i)) { - continue; - } - - int value = column.getInt(i); - updateStateByRemove(value); - } - } - - private void removeLongInput(Column column) { - for (int i = 0; i < column.getPositionCount(); i++) { - if (column.isNull(i)) { - continue; - } - - long value = column.getLong(i); - updateStateByRemove(value); - } - } - - private void removeFloatInput(Column column) { - for (int i = 0; i < column.getPositionCount(); i++) { - if (column.isNull(i)) { - continue; - } - - float value = column.getFloat(i); - updateStateByRemove(value); - } - } - - private void removeDoubleInput(Column column) { - for (int i = 0; i < column.getPositionCount(); i++) { - if (column.isNull(i)) { - continue; - } - - double value = column.getDouble(i); - updateStateByRemove(value); - } - } - private void updateStateByRemove(double value) { long newCount = count - 1; double newMean = (count * mean - value) / newCount; @@ -407,4 +233,15 @@ public class TableVarianceAccumulator implements TableAccumulator { count = newCount; mean = newMean; } + + private void checkInputDataType() { + if (!seriesDataType.isNumeric()) { + throw unsupportedDataTypeException(); + } + } + + private UnSupportedDataTypeException unsupportedDataTypeException() { + return new UnSupportedDataTypeException( + String.format(CalcMessages.UNSUPPORTED_DATA_TYPE_IN_VARIANCE_AGGREGATION, seriesDataType)); + } } diff --git a/iotdb-core/calc-commons/src/main/java/org/apache/iotdb/calc/execution/operator/source/relational/aggregation/grouped/GroupedVarianceAccumulator.java b/iotdb-core/calc-commons/src/main/java/org/apache/iotdb/calc/execution/operator/source/relational/aggregation/grouped/GroupedVarianceAccumulator.java index 75446b2cfab..9c70b0078d6 100644 --- a/iotdb-core/calc-commons/src/main/java/org/apache/iotdb/calc/execution/operator/source/relational/aggregation/grouped/GroupedVarianceAccumulator.java +++ b/iotdb-core/calc-commons/src/main/java/org/apache/iotdb/calc/execution/operator/source/relational/aggregation/grouped/GroupedVarianceAccumulator.java @@ -23,6 +23,8 @@ import org.apache.iotdb.calc.execution.aggregation.VarianceAccumulator; import org.apache.iotdb.calc.execution.operator.source.relational.aggregation.AggregationMask; import org.apache.iotdb.calc.execution.operator.source.relational.aggregation.grouped.array.DoubleBigArray; import org.apache.iotdb.calc.execution.operator.source.relational.aggregation.grouped.array.LongBigArray; +import org.apache.iotdb.calc.i18n.CalcMessages; +import org.apache.iotdb.calc.utils.TypeServices; import org.apache.tsfile.block.column.Column; import org.apache.tsfile.block.column.ColumnBuilder; @@ -30,6 +32,7 @@ import org.apache.tsfile.enums.TSDataType; import org.apache.tsfile.read.common.block.column.BinaryColumn; import org.apache.tsfile.read.common.block.column.BinaryColumnBuilder; import org.apache.tsfile.read.common.block.column.RunLengthEncodedColumn; +import org.apache.tsfile.read.common.type.Type; import org.apache.tsfile.utils.Binary; import org.apache.tsfile.utils.BytesUtils; import org.apache.tsfile.utils.RamUsageEstimator; @@ -42,6 +45,7 @@ public class GroupedVarianceAccumulator implements GroupedAccumulator { private static final long INSTANCE_SIZE = RamUsageEstimator.shallowSizeOfInstance(GroupedVarianceAccumulator.class); private final TSDataType seriesDataType; + private final TypeServices.ColumnToDoubleConverter doubleValueConverter; private final VarianceAccumulator.VarianceType varianceType; private final LongBigArray counts = new LongBigArray(); @@ -51,6 +55,10 @@ public class GroupedVarianceAccumulator implements GroupedAccumulator { public GroupedVarianceAccumulator( TSDataType seriesDataType, VarianceAccumulator.VarianceType varianceType) { this.seriesDataType = seriesDataType; + this.doubleValueConverter = + TypeServices.NUMERIC_COLUMN_TO_DOUBLE_CONVERTER_SERVICE + .call(Type.fromTsDataType(seriesDataType)) + .create(this::unsupportedDataTypeException); this.varianceType = varianceType; } @@ -68,30 +76,8 @@ public class GroupedVarianceAccumulator implements GroupedAccumulator { @Override public void addInput(int[] groupIds, Column[] arguments, AggregationMask mask) { - switch (seriesDataType) { - case INT32: - addIntInput(groupIds, arguments[0], mask); - return; - case INT64: - addLongInput(groupIds, arguments[0], mask); - return; - case FLOAT: - addFloatInput(groupIds, arguments[0], mask); - return; - case DOUBLE: - addDoubleInput(groupIds, arguments[0], mask); - return; - case TEXT: - case BLOB: - case OBJECT: - case BOOLEAN: - case DATE: - case STRING: - case TIMESTAMP: - default: - throw new UnSupportedDataTypeException( - String.format("Unsupported data type in VARIANCE Aggregation: %s", seriesDataType)); - } + checkInputDataType(); + updateStateByAdd(groupIds, arguments[0], mask); } @Override @@ -191,7 +177,7 @@ public class GroupedVarianceAccumulator implements GroupedAccumulator { m2s.reset(); } - private void addIntInput(int[] groupIds, Column column, AggregationMask mask) { + private void updateStateByAdd(int[] groupIds, Column column, AggregationMask mask) { int positionCount = mask.getSelectedPositionCount(); if (mask.isSelectAll()) { @@ -200,7 +186,7 @@ public class GroupedVarianceAccumulator implements GroupedAccumulator { continue; } - int value = column.getInt(i); + double value = doubleValueConverter.convert(column, i); counts.increment(groupIds[i]); double delta = value - means.get(groupIds[i]); means.add(groupIds[i], delta / counts.get(groupIds[i])); @@ -215,7 +201,7 @@ public class GroupedVarianceAccumulator implements GroupedAccumulator { continue; } - int value = column.getInt(position); + double value = doubleValueConverter.convert(column, position); counts.increment(groupIds[position]); double delta = value - means.get(groupIds[position]); means.add(groupIds[position], delta / counts.get(groupIds[position])); @@ -224,102 +210,14 @@ public class GroupedVarianceAccumulator implements GroupedAccumulator { } } - private void addLongInput(int[] groupIds, Column column, AggregationMask mask) { - int positionCount = mask.getSelectedPositionCount(); - - if (mask.isSelectAll()) { - for (int i = 0; i < positionCount; i++) { - if (column.isNull(i)) { - continue; - } - - long value = column.getLong(i); - counts.increment(groupIds[i]); - double delta = value - means.get(groupIds[i]); - means.add(groupIds[i], delta / counts.get(groupIds[i])); - m2s.add(groupIds[i], delta * (value - means.get(groupIds[i]))); - } - } else { - int[] selectedPositions = mask.getSelectedPositions(); - int position; - for (int i = 0; i < positionCount; i++) { - position = selectedPositions[i]; - if (column.isNull(position)) { - continue; - } - - long value = column.getLong(position); - counts.increment(groupIds[position]); - double delta = value - means.get(groupIds[position]); - means.add(groupIds[position], delta / counts.get(groupIds[position])); - m2s.add(groupIds[position], delta * (value - means.get(groupIds[position]))); - } + private void checkInputDataType() { + if (!seriesDataType.isNumeric()) { + throw unsupportedDataTypeException(); } } - private void addFloatInput(int[] groupIds, Column column, AggregationMask mask) { - int positionCount = mask.getSelectedPositionCount(); - - if (mask.isSelectAll()) { - for (int i = 0; i < positionCount; i++) { - if (column.isNull(i)) { - continue; - } - - float value = column.getFloat(i); - counts.increment(groupIds[i]); - double delta = value - means.get(groupIds[i]); - means.add(groupIds[i], delta / counts.get(groupIds[i])); - m2s.add(groupIds[i], delta * (value - means.get(groupIds[i]))); - } - } else { - int[] selectedPositions = mask.getSelectedPositions(); - int position; - for (int i = 0; i < positionCount; i++) { - position = selectedPositions[i]; - if (column.isNull(position)) { - continue; - } - - float value = column.getFloat(position); - counts.increment(groupIds[position]); - double delta = value - means.get(groupIds[position]); - means.add(groupIds[position], delta / counts.get(groupIds[position])); - m2s.add(groupIds[position], delta * (value - means.get(groupIds[position]))); - } - } - } - - private void addDoubleInput(int[] groupIds, Column column, AggregationMask mask) { - int positionCount = mask.getSelectedPositionCount(); - - if (mask.isSelectAll()) { - for (int i = 0; i < positionCount; i++) { - if (column.isNull(i)) { - continue; - } - - double value = column.getDouble(i); - counts.increment(groupIds[i]); - double delta = value - means.get(groupIds[i]); - means.add(groupIds[i], delta / counts.get(groupIds[i])); - m2s.add(groupIds[i], delta * (value - means.get(groupIds[i]))); - } - } else { - int[] selectedPositions = mask.getSelectedPositions(); - int position; - for (int i = 0; i < positionCount; i++) { - position = selectedPositions[i]; - if (column.isNull(position)) { - continue; - } - - double value = column.getDouble(position); - counts.increment(groupIds[position]); - double delta = value - means.get(groupIds[position]); - means.add(groupIds[position], delta / counts.get(groupIds[position])); - m2s.add(groupIds[position], delta * (value - means.get(groupIds[position]))); - } - } + private UnSupportedDataTypeException unsupportedDataTypeException() { + return new UnSupportedDataTypeException( + String.format(CalcMessages.UNSUPPORTED_DATA_TYPE_IN_VARIANCE_AGGREGATION, seriesDataType)); } } diff --git a/iotdb-core/calc-commons/src/test/java/org/apache/iotdb/calc/execution/operator/source/relational/aggregation/VarianceAccumulatorTest.java b/iotdb-core/calc-commons/src/test/java/org/apache/iotdb/calc/execution/operator/source/relational/aggregation/VarianceAccumulatorTest.java new file mode 100644 index 00000000000..872eeaa2a33 --- /dev/null +++ b/iotdb-core/calc-commons/src/test/java/org/apache/iotdb/calc/execution/operator/source/relational/aggregation/VarianceAccumulatorTest.java @@ -0,0 +1,126 @@ +/* + * 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. + */ + +package org.apache.iotdb.calc.execution.operator.source.relational.aggregation; + +import org.apache.iotdb.calc.execution.aggregation.VarianceAccumulator; +import org.apache.iotdb.calc.execution.operator.source.relational.aggregation.grouped.GroupedVarianceAccumulator; + +import org.apache.tsfile.block.column.Column; +import org.apache.tsfile.block.column.ColumnBuilder; +import org.apache.tsfile.enums.TSDataType; +import org.apache.tsfile.read.common.block.column.DoubleColumnBuilder; +import org.apache.tsfile.read.common.block.column.DoubleColumn; +import org.apache.tsfile.read.common.block.column.FloatColumn; +import org.apache.tsfile.read.common.block.column.IntColumn; +import org.apache.tsfile.read.common.block.column.LongColumn; +import org.apache.tsfile.write.UnSupportedDataTypeException; +import org.junit.Assert; +import org.junit.Test; + +import java.util.Optional; + +public class VarianceAccumulatorTest { + + @Test + public void testVarianceAccumulatorsReadAllNumericTypes() { + TSDataType[] dataTypes = { + TSDataType.INT32, TSDataType.INT64, TSDataType.FLOAT, TSDataType.DOUBLE + }; + Column[] valueColumns = { + new IntColumn(2, Optional.empty(), new int[] {1, 3}), + new LongColumn(2, Optional.empty(), new long[] {1, 3}), + new FloatColumn(2, Optional.empty(), new float[] {1, 3}), + new DoubleColumn(2, Optional.empty(), new double[] {1, 3}) + }; + for (int i = 0; i < dataTypes.length; i++) { + TSDataType dataType = dataTypes[i]; + Column valueColumn = valueColumns[i]; + + VarianceAccumulator treeAccumulator = + new VarianceAccumulator(dataType, VarianceAccumulator.VarianceType.VAR_POP); + treeAccumulator.addInput(new Column[] {valueColumn, valueColumn}, /* bitMap= */ null); + DoubleColumnBuilder treeResult = new DoubleColumnBuilder(null, 1); + treeAccumulator.outputFinal(treeResult); + Assert.assertEquals(1.0, treeResult.build().getDouble(0), 0.0); + + TableVarianceAccumulator tableAccumulator = + new TableVarianceAccumulator(dataType, VarianceAccumulator.VarianceType.VAR_POP); + tableAccumulator.addInput( + new Column[] {valueColumn}, + AggregationMask.createSelectAll(valueColumn.getPositionCount())); + DoubleColumnBuilder tableResult = new DoubleColumnBuilder(null, 1); + tableAccumulator.evaluateFinal(tableResult); + Assert.assertEquals(1.0, tableResult.build().getDouble(0), 0.0); + + GroupedVarianceAccumulator groupedAccumulator = + new GroupedVarianceAccumulator(dataType, VarianceAccumulator.VarianceType.VAR_POP); + groupedAccumulator.setGroupCount(1); + groupedAccumulator.addInput( + new int[] {0, 0}, + new Column[] {valueColumn}, + AggregationMask.createSelectAll(valueColumn.getPositionCount())); + DoubleColumnBuilder groupedResult = new DoubleColumnBuilder(null, 1); + groupedAccumulator.evaluateFinal(0, groupedResult); + Assert.assertEquals(1.0, groupedResult.build().getDouble(0), 0.0); + } + } + + @Test + public void testVarianceAccumulatorsStillRejectTemporalTypes() { + TSDataType[] dataTypes = {TSDataType.DATE, TSDataType.TIMESTAMP}; + Column[] valueColumns = { + new IntColumn(2, Optional.empty(), new int[] {1, 3}, TSDataType.DATE), + new LongColumn(2, Optional.empty(), new long[] {1, 3}) + }; + for (int i = 0; i < dataTypes.length; i++) { + TSDataType dataType = dataTypes[i]; + Column valueColumn = valueColumns[i]; + + VarianceAccumulator treeAccumulator = + new VarianceAccumulator(dataType, VarianceAccumulator.VarianceType.VAR_POP); + Assert.assertThrows( + UnSupportedDataTypeException.class, + () -> + treeAccumulator.addInput( + new Column[] {valueColumn, valueColumn}, /* bitMap= */ null)); + + TableVarianceAccumulator tableAccumulator = + new TableVarianceAccumulator(dataType, VarianceAccumulator.VarianceType.VAR_POP); + Assert.assertThrows( + UnSupportedDataTypeException.class, + () -> + tableAccumulator.addInput( + new Column[] {valueColumn}, + AggregationMask.createSelectAll(valueColumn.getPositionCount()))); + Assert.assertThrows( + UnSupportedDataTypeException.class, + () -> tableAccumulator.removeInput(new Column[] {valueColumn})); + + GroupedVarianceAccumulator groupedAccumulator = + new GroupedVarianceAccumulator(dataType, VarianceAccumulator.VarianceType.VAR_POP); + Assert.assertThrows( + UnSupportedDataTypeException.class, + () -> + groupedAccumulator.addInput( + new int[] {0, 0}, + new Column[] {valueColumn}, + AggregationMask.createSelectAll(valueColumn.getPositionCount()))); + } + } +}
