This is an automated email from the ASF dual-hosted git repository. ycycse pushed a commit to branch checkBugInArima in repository https://gitbox.apache.org/repos/asf/iotdb.git
commit 1ab8f5c740a0ec3b9e6356bdff42bab76e524639 Author: YangCaiyin <[email protected]> AuthorDate: Mon May 26 20:55:15 2025 +0800 support holtwinters and enhance codes --- .../org/apache/iotdb/ainode/it/AINodeBasicIT.java | 20 ++++++++++++++++++++ iotdb-core/ainode/ainode/core/constant.py | 1 + .../ainode/core/model/built_in_model_factory.py | 4 ++-- .../iotdb/confignode/persistence/ModelInfo.java | 3 ++- .../function/tvf/ForecastTableFunction.java | 2 +- 5 files changed, 26 insertions(+), 4 deletions(-) diff --git a/integration-test/src/test/java/org/apache/iotdb/ainode/it/AINodeBasicIT.java b/integration-test/src/test/java/org/apache/iotdb/ainode/it/AINodeBasicIT.java index 07d29c0d224..40aad154f56 100644 --- a/integration-test/src/test/java/org/apache/iotdb/ainode/it/AINodeBasicIT.java +++ b/integration-test/src/test/java/org/apache/iotdb/ainode/it/AINodeBasicIT.java @@ -176,6 +176,26 @@ public class AINodeBasicIT { } } + @Test + public void callInferenceTest2() { + String sql = + "CALL INFERENCE(_holtwinters, \"select s0 from root.AI.data\", predict_length=6, generateTime=true)"; + try (Connection connection = EnvFactory.getEnv().getConnection(); + Statement statement = connection.createStatement()) { + try (ResultSet resultSet = statement.executeQuery(sql)) { + ResultSetMetaData resultSetMetaData = resultSet.getMetaData(); + checkHeader(resultSetMetaData, "Time,output0"); + int count = 0; + while (resultSet.next()) { + count++; + } + assertEquals(6, count); + } + }catch (SQLException e) { + fail(e.getMessage()); + } + } + @Test public void callInferenceTest() { String sql = diff --git a/iotdb-core/ainode/ainode/core/constant.py b/iotdb-core/ainode/ainode/core/constant.py index a80ca680989..eb7f7b717bd 100644 --- a/iotdb-core/ainode/ainode/core/constant.py +++ b/iotdb-core/ainode/ainode/core/constant.py @@ -139,6 +139,7 @@ class ModelInputName(Enum): class BuiltInModelType(Enum): # forecast models ARIMA = "_arima" + HOLTWINTERS = "_hotwinters" EXPONENTIAL_SMOOTHING = "_exponentialsmoothing" NAIVE_FORECASTER = "_naiveforecaster" STL_FORECASTER = "_stlforecaster" diff --git a/iotdb-core/ainode/ainode/core/model/built_in_model_factory.py b/iotdb-core/ainode/ainode/core/model/built_in_model_factory.py index 0d2991a7f9f..662523f0ce9 100644 --- a/iotdb-core/ainode/ainode/core/model/built_in_model_factory.py +++ b/iotdb-core/ainode/ainode/core/model/built_in_model_factory.py @@ -85,7 +85,7 @@ def fetch_built_in_model(model_id, inference_attributes): # build the built-in model if model_id == BuiltInModelType.ARIMA.value: model = ArimaModel(attributes) - elif model_id == BuiltInModelType.EXPONENTIAL_SMOOTHING.value: + elif model_id == BuiltInModelType.EXPONENTIAL_SMOOTHING.value or model_id == BuiltInModelType.HOLTWINTERS.value: model = ExponentialSmoothingModel(attributes) elif model_id == BuiltInModelType.NAIVE_FORECASTER.value: model = NaiveForecasterModel(attributes) @@ -460,7 +460,7 @@ arima_attribute_map = { ), AttributeName.ORDER.value: TupleAttribute( name=AttributeName.ORDER.value, - default_value=(1, 0, 0), + default_value=(96, 1, 96), value_type=int ), AttributeName.SEASONAL_ORDER.value: TupleAttribute( diff --git a/iotdb-core/confignode/src/main/java/org/apache/iotdb/confignode/persistence/ModelInfo.java b/iotdb-core/confignode/src/main/java/org/apache/iotdb/confignode/persistence/ModelInfo.java index e2beede330a..a72873e9aa7 100644 --- a/iotdb-core/confignode/src/main/java/org/apache/iotdb/confignode/persistence/ModelInfo.java +++ b/iotdb-core/confignode/src/main/java/org/apache/iotdb/confignode/persistence/ModelInfo.java @@ -75,10 +75,11 @@ public class ModelInfo implements SnapshotProcessor { private static final Set<String> builtInAnomalyDetectionModel = new HashSet<>(); static { - builtInForecastModel.add("_timerxl"); + builtInForecastModel.add("_TimerXL"); builtInForecastModel.add("_ARIMA"); builtInForecastModel.add("_NaiveForecaster"); builtInForecastModel.add("_STLForecaster"); + builtInForecastModel.add("_HoltWinters"); builtInForecastModel.add("_ExponentialSmoothing"); builtInAnomalyDetectionModel.add("_GaussianHMM"); builtInAnomalyDetectionModel.add("_GMMHMM"); diff --git a/iotdb-core/datanode/src/main/java/org/apache/iotdb/db/queryengine/plan/relational/function/tvf/ForecastTableFunction.java b/iotdb-core/datanode/src/main/java/org/apache/iotdb/db/queryengine/plan/relational/function/tvf/ForecastTableFunction.java index dbe91526bd8..69d637fae52 100644 --- a/iotdb-core/datanode/src/main/java/org/apache/iotdb/db/queryengine/plan/relational/function/tvf/ForecastTableFunction.java +++ b/iotdb-core/datanode/src/main/java/org/apache/iotdb/db/queryengine/plan/relational/function/tvf/ForecastTableFunction.java @@ -333,7 +333,7 @@ public class ForecastTableFunction implements TableFunction { } } } else { - String[] predictedColumnsArray = predicatedColumns.split(","); + String[] predictedColumnsArray = predicatedColumns.split(";"); Map<String, Integer> inputColumnIndexMap = new HashMap<>(); for (int i = 0, size = allInputColumnsName.size(); i < size; i++) { Optional<String> fieldName = allInputColumnsName.get(i);
