Java时序预测案例

wen java案例 3

本文目录导读:

Java时序预测案例

  1. 项目依赖
  2. 时序数据模型
  3. 数据预处理工具类
  4. ARIMA模型实现
  5. 时序预测运行示例
  6. 使用示例

我来提供一个Java时序预测的完整案例,使用ARIMA模型进行预测。

项目依赖

<!-- pom.xml -->
<dependencies>
    <!-- Apache Commons Math for 统计计算 -->
    <dependency>
        <groupId>org.apache.commons</groupId>
        <artifactId>commons-math3</artifactId>
        <version>3.6.1</version>
    </dependency>
    <!-- Lombok 简化代码 -->
    <dependency>
        <groupId>org.projectlombok</groupId>
        <artifactId>lombok</artifactId>
        <version>1.18.30</version>
        <optional>true</optional>
    </dependency>
    <!-- CSV读取 -->
    <dependency>
        <groupId>com.opencsv</groupId>
        <artifactId>opencsv</artifactId>
        <version>5.8</version>
    </dependency>
</dependencies>

时序数据模型

import lombok.AllArgsConstructor;
import lombok.Data;
import lombok.NoArgsConstructor;
import java.time.LocalDateTime;
@Data
@NoArgsConstructor
@AllArgsConstructor
public class TimeSeriesPoint {
    private LocalDateTime timestamp;
    private double value;
}
@Data
@NoArgsConstructor
@AllArgsConstructor
public class ForecastResult {
    private double[] predictedValues;
    private double[] confidenceUpper;
    private double[] confidenceLower;
    private double rmse;
    private double mape;
}

数据预处理工具类

import java.util.*;
import java.util.stream.Collectors;
public class DataPreprocessor {
    // 处理缺失值 - 线性插值
    public static double[] handleMissingValues(double[] data) {
        double[] result = data.clone();
        for (int i = 0; i < result.length; i++) {
            if (Double.isNaN(result[i])) {
                // 找到前后有效值
                double prevValue = findPrevValid(result, i);
                double nextValue = findNextValid(result, i);
                int prevIndex = findPrevValidIndex(result, i);
                int nextIndex = findNextValidIndex(result, i);
                if (prevIndex == -1 && nextIndex == -1) {
                    result[i] = 0;
                } else if (prevIndex == -1) {
                    result[i] = nextValue;
                } else if (nextIndex == -1) {
                    result[i] = prevValue;
                } else {
                    // 线性插值
                    double ratio = (double)(i - prevIndex) / (nextIndex - prevIndex);
                    result[i] = prevValue + ratio * (nextValue - prevValue);
                }
            }
        }
        return result;
    }
    // 数据标准化 (Z-Score归一化)
    public static double[] normalize(double[] data) {
        double mean = Arrays.stream(data).average().orElse(0);
        double std = calculateStd(data, mean);
        return Arrays.stream(data)
                .map(x -> (x - mean) / std)
                .toArray();
    }
    // 数据缩放 (Min-Max归一化)
    public static double[] scale(double[] data, double min, double max) {
        double dataMin = Arrays.stream(data).min().orElse(0);
        double dataMax = Arrays.stream(data).max().orElse(1);
        return Arrays.stream(data)
                .map(x -> min + (x - dataMin) * (max - min) / (dataMax - dataMin))
                .toArray();
    }
    // 差分处理 (用于非平稳序列)
    public static double[] difference(double[] data, int lag) {
        double[] result = new double[data.length - lag];
        for (int i = lag; i < data.length; i++) {
            result[i - lag] = data[i] - data[i - lag];
        }
        return result;
    }
    // 逆差分
    public static double[] inverseDifference(double[] original, double[] diffData, int lag) {
        double[] result = new double[original.length];
        System.arraycopy(original, 0, result, 0, lag);
        for (int i = 0; i < diffData.length; i++) {
            result[i + lag] = result[i + lag - 1] + diffData[i];
        }
        return result;
    }
    private static double findPrevValid(double[] data, int index) {
        for (int i = index - 1; i >= 0; i--) {
            if (!Double.isNaN(data[i])) return data[i];
        }
        return Double.NaN;
    }
    private static int findPrevValidIndex(double[] data, int index) {
        for (int i = index - 1; i >= 0; i--) {
            if (!Double.isNaN(data[i])) return i;
        }
        return -1;
    }
    private static double findNextValid(double[] data, int index) {
        for (int i = index + 1; i < data.length; i++) {
            if (!Double.isNaN(data[i])) return data[i];
        }
        return Double.NaN;
    }
    private static int findNextValidIndex(double[] data, int index) {
        for (int i = index + 1; i < data.length; i++) {
            if (!Double.isNaN(data[i])) return i;
        }
        return -1;
    }
    private static double calculateStd(double[] data, double mean) {
        double sum = 0;
        for (double d : data) {
            sum += Math.pow(d - mean, 2);
        }
        return Math.sqrt(sum / (data.length - 1));
    }
}

ARIMA模型实现

import org.apache.commons.math3.linear.*;
import org.apache.commons.math3.stat.regression.OLSMultipleLinearRegression;
import java.util.*;
public class ARIMAModel {
    private int p; // 自回归阶数
    private int d; // 差分阶数
    private int q; // 移动平均阶数
    private double[] arCoefficients;
    private double[] maCoefficients;
    private double[] residuals;
    private double[] originalData;
    private double[] differencedData;
    private double mean;
    public ARIMAModel(int p, int d, int q) {
        this.p = p;
        this.d = d;
        this.q = q;
    }
    // 训练模型
    public void fit(double[] data) {
        this.originalData = data.clone();
        // 进行差分
        double[] temp = data.clone();
        for (int i = 0; i < d; i++) {
            temp = DataPreprocessor.difference(temp, 1);
        }
        this.differencedData = temp;
        // 计算均值
        this.mean = Arrays.stream(differencedData).average().orElse(0);
        // 中心化
        double[] centeredData = Arrays.stream(differencedData)
                .map(x -> x - mean)
                .toArray();
        // 估计AR和MA参数
        estimateParameters(centeredData);
    }
    // 预测
    public double[] predict(int steps) {
        double[] predictions = new double[steps];
        double[] lastValues = getLastValues(p);
        double[] lastResiduals = getLastResiduals(q);
        for (int i = 0; i < steps; i++) {
            double prediction = mean;
            // AR部分
            if (arCoefficients != null) {
                for (int j = 0; j < p && j < lastValues.length; j++) {
                    prediction += arCoefficients[j] * (lastValues[lastValues.length - 1 - j] - mean);
                }
            }
            // MA部分
            if (maCoefficients != null && lastResiduals != null) {
                for (int j = 0; j < q && j < lastResiduals.length; j++) {
                    prediction += maCoefficients[j] * lastResiduals[lastResiduals.length - 1 - j];
                }
            }
            predictions[i] = prediction;
            // 更新最后值
            double[] newLastValues = new double[lastValues.length + 1];
            System.arraycopy(lastValues, 0, newLastValues, 0, lastValues.length);
            newLastValues[lastValues.length] = prediction;
            lastValues = newLastValues;
            // 更新残差
            if (lastResiduals != null) {
                double[] newLastResiduals = new double[lastResiduals.length + 1];
                System.arraycopy(lastResiduals, 0, newLastResiduals, 0, lastResiduals.length);
                newLastResiduals[lastResiduals.length] = 0;
                lastResiduals = newLastResiduals;
            }
        }
        // 逆差分还原预测值
        if (d > 0) {
            predictions = DataPreprocessor.inverseDifference(
                    originalData, predictions, d);
        }
        return predictions;
    }
    // 评估模型
    public ForecastResult evaluate(double[] testData) {
        int steps = testData.length;
        double[] predictions = predict(steps);
        // 计算RMSE
        double sumSquaredError = 0;
        double sumAbsPercentageError = 0;
        for (int i = 0; i < steps; i++) {
            double error = testData[i] - predictions[i];
            sumSquaredError += error * error;
            if (testData[i] != 0) {
                sumAbsPercentageError += Math.abs(error / testData[i]);
            }
        }
        double rmse = Math.sqrt(sumSquaredError / steps);
        double mape = (sumAbsPercentageError / steps) * 100;
        // 计算置信区间
        double stdError = calculateStandardError();
        double[] confidenceUpper = new double[steps];
        double[] confidenceLower = new double[steps];
        for (int i = 0; i < steps; i++) {
            confidenceUpper[i] = predictions[i] + 1.96 * stdError * Math.sqrt(i + 1);
            confidenceLower[i] = predictions[i] - 1.96 * stdError * Math.sqrt(i + 1);
        }
        return new ForecastResult(predictions, confidenceUpper, confidenceLower, rmse, mape);
    }
    private void estimateParameters(double[] data) {
        int n = data.length;
        int maxLag = Math.max(p, q);
        if (n <= maxLag) {
            return;
        }
        // 使用Yule-Walker方程估计AR参数
        if (p > 0) {
            double[] acf = calculateACF(data, p);
            arCoefficients = solveYuleWalker(acf);
        }
        // 计算残差
        residuals = new double[n - maxLag];
        for (int i = maxLag; i < n; i++) {
            double arPart = mean;
            if (arCoefficients != null) {
                for (int j = 0; j < p; j++) {
                    arPart += arCoefficients[j] * (data[i - 1 - j] - mean);
                }
            }
            residuals[i - maxLag] = data[i] - arPart;
        }
        // 估计MA参数 (简化版)
        if (q > 0 && residuals.length > q) {
            maCoefficients = estimateMA(residuals);
        }
    }
    private double[] calculateACF(double[] data, int maxLag) {
        int n = data.length;
        double mean = Arrays.stream(data).average().orElse(0);
        double variance = 0;
        for (double d : data) {
            variance += Math.pow(d - mean, 2);
        }
        variance /= n;
        double[] acf = new double[maxLag + 1];
        acf[0] = 1.0;
        for (int k = 1; k <= maxLag; k++) {
            double sum = 0;
            for (int i = k; i < n; i++) {
                sum += (data[i] - mean) * (data[i - k] - mean);
            }
            acf[k] = sum / (n * variance);
        }
        return acf;
    }
    private double[] solveYuleWalker(double[] acf) {
        int m = p;
        RealMatrix R = new Array2DRowRealMatrix(m, m);
        RealVector r = new ArrayRealVector(m);
        for (int i = 0; i < m; i++) {
            for (int j = 0; j < m; j++) {
                R.setEntry(i, j, acf[Math.abs(i - j)]);
            }
            r.setEntry(i, acf[i + 1]);
        }
        try {
            DecompositionSolver solver = new LUDecomposition(R).getSolver();
            RealVector phi = solver.solve(r);
            return phi.toArray();
        } catch (SingularMatrixException e) {
            return new double[m];
        }
    }
    private double[] estimateMA(double[] residuals) {
        int n = residuals.length;
        OLSMultipleLinearRegression regression = new OLSMultipleLinearRegression();
        double[][] x = new double[n - q][q];
        double[] y = new double[n - q];
        for (int i = q; i < n; i++) {
            for (int j = 0; j < q; j++) {
                x[i - q][j] = residuals[i - 1 - j];
            }
            y[i - q] = residuals[i];
        }
        try {
            regression.newSampleData(y, x);
            return regression.estimateRegressionParameters();
        } catch (Exception e) {
            return new double[q];
        }
    }
    private double[] getLastValues(int count) {
        if (differencedData == null || differencedData.length == 0) {
            return new double[0];
        }
        int n = differencedData.length;
        int size = Math.min(count, n);
        double[] result = new double[size];
        for (int i = 0; i < size; i++) {
            result[size - 1 - i] = differencedData[n - 1 - i];
        }
        return result;
    }
    private double[] getLastResiduals(int count) {
        if (residuals == null || residuals.length == 0) {
            return null;
        }
        int n = residuals.length;
        int size = Math.min(count, n);
        double[] result = new double[size];
        for (int i = 0; i < size; i++) {
            result[size - 1 - i] = residuals[n - 1 - i];
        }
        return result;
    }
    private double calculateStandardError() {
        if (residuals == null || residuals.length == 0) {
            return 0;
        }
        double sum = 0;
        for (double r : residuals) {
            sum += r * r;
        }
        return Math.sqrt(sum / (residuals.length - p - q));
    }
    // 获取模型信息
    public String getModelInfo() {
        StringBuilder sb = new StringBuilder();
        sb.append(String.format("ARIMA(%d,%d,%d) Model\n", p, d, q));
        if (arCoefficients != null) {
            sb.append("AR Coefficients: ");
            sb.append(Arrays.toString(arCoefficients));
            sb.append("\n");
        }
        if (maCoefficients != null) {
            sb.append("MA Coefficients: ");
            sb.append(Arrays.toString(maCoefficients));
            sb.append("\n");
        }
        sb.append(String.format("Mean: %.4f\n", mean));
        return sb.toString();
    }
}

时序预测运行示例

import java.time.LocalDateTime;
import java.util.*;
public class TimeSeriesForecastExample {
    public static void main(String[] args) {
        // 生成示例数据
        double[] data = generateSampleData(100);
        // 数据预处理
        DataPreprocessor preprocessor = new DataPreprocessor();
        double[] cleanedData = preprocessor.handleMissingValues(data);
        // 划分训练集和测试集
        int trainSize = 80;
        int testSize = 20;
        double[] trainData = Arrays.copyOfRange(cleanedData, 0, trainSize);
        double[] testData = Arrays.copyOfRange(cleanedData, trainSize, trainSize + testSize);
        // 创建并训练ARIMA模型
        ARIMAModel model = new ARIMAModel(2, 1, 2);
        model.fit(trainData);
        System.out.println("========== 模型信息 ==========");
        System.out.println(model.getModelInfo());
        // 预测
        double[] predictions = model.predict(testSize);
        // 评估
        ForecastResult result = model.evaluate(testData);
        System.out.println("\n========== 预测结果 ==========");
        for (int i = 0; i < testSize; i++) {
            System.out.printf("时间 %d: 实际值=%.4f, 预测值=%.4f, 置信区间=[%.4f, %.4f]\n",
                    i + 1,
                    testData[i],
                    result.getPredictedValues()[i],
                    result.getConfidenceLower()[i],
                    result.getConfidenceUpper()[i]);
        }
        System.out.println("\n========== 模型评估 ==========");
        System.out.printf("RMSE: %.4f\n", result.getRmse());
        System.out.printf("MAPE: %.2f%%\n", result.getMape());
        // 可视化数据预览
        System.out.println("\n========== 数据预览 ==========");
        System.out.println("原始数据:");
        for (int i = 0; i < 10; i++) {
            System.out.printf("%.2f ", data[i]);
        }
        System.out.println("...");
        // 示例:批量预测多个时段
        System.out.println("\n========== 未来10期预测 ==========");
        int futureSteps = 10;
        double[] futurePredictions = model.predict(testSize + futureSteps);
        for (int i = 0; i < futureSteps; i++) {
            System.out.printf("第%d期: %.4f\n", i + 1, futurePredictions[trainSize + testSize + i]);
        }
    }
    // 生成示例数据 (带有趋势和季节性的时间序列)
    private static double[] generateSampleData(int n) {
        Random random = new Random(42);
        double[] data = new double[n];
        double trend = 0.1;
        double seasonalAmplitude = 5.0;
        int seasonalPeriod = 12;
        double noise = 2.0;
        double base = 100;
        for (int i = 0; i < n; i++) {
            // 趋势
            double trendComponent = trend * i;
            // 季节性
            double seasonalComponent = seasonalAmplitude * 
                    Math.sin(2 * Math.PI * i / seasonalPeriod);
            // 随机噪声
            double noiseComponent = noise * random.nextGaussian();
            data[i] = base + trendComponent + seasonalComponent + noiseComponent;
        }
        return data;
    }
}

使用示例

public class SimpleUsageExample {
    public static void main(String[] args) {
        // 准备数据
        double[] salesData = {120, 135, 128, 142, 145, 138, 
                             152, 148, 155, 160, 158, 165};
        // 创建模型 (ARIMA(1,1,1))
        ARIMAModel model = new ARIMAModel(1, 1, 1);
        // 训练模型
        model.fit(salesData);
        // 预测未来3个月
        double[] forecast = model.predict(3);
        System.out.println("销量预测:");
        for (int i = 0; i < forecast.length; i++) {
            System.out.printf("Month %d: %.2f\n", i + 13, forecast[i]);
        }
    }
}

这个时序预测案例提供了:

  1. 完整的数据预处理:缺失值处理、标准化、差分等
  2. ARIMA模型实现:自回归(AR)、差分(I)、移动平均(MA)组件
  3. 模型评估:RMSE、MAPE等指标
  4. 置信区间计算:预测的不确定性量化
  5. 示例数据生成:包含趋势、季节性和噪声

可以根据具体业务需求调整模型参数(p, d, q)来优化预测效果。

上一篇AIOps案例

下一篇Java聚类案例

抱歉,评论功能暂时关闭!