本文目录导读:

我来提供一个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]);
}
}
}
这个时序预测案例提供了:
- 完整的数据预处理:缺失值处理、标准化、差分等
- ARIMA模型实现:自回归(AR)、差分(I)、移动平均(MA)组件
- 模型评估:RMSE、MAPE等指标
- 置信区间计算:预测的不确定性量化
- 示例数据生成:包含趋势、季节性和噪声
可以根据具体业务需求调整模型参数(p, d, q)来优化预测效果。