package cn.iocoder.yudao.module.bi.service.decision;
|
|
import cn.hutool.core.collection.CollUtil;
|
|
import java.math.BigDecimal;
|
import java.math.RoundingMode;
|
import java.util.ArrayList;
|
import java.util.List;
|
|
/**
|
* Java 统计预测引擎
|
*
|
* 提供 SMA/WMA/线性回归/季节指数/同比环比外推 等统计预测算法,
|
* 输入历史序列,返回预测值及 95% 置信区间。
|
*
|
* @author 超级管理员
|
*/
|
public final class BiForecastEngine {
|
|
private BiForecastEngine() {
|
}
|
|
/**
|
* 预测结果:值 + 置信区间
|
*/
|
public record Forecast(BigDecimal value, BigDecimal lowerBound, BigDecimal upperBound) {
|
}
|
|
/**
|
* 统一预测入口
|
*
|
* @param model SMA / WMA / LR / SEASONAL / YOY
|
* @param history 历史序列(按时间升序)
|
* @param window 移动平均窗口(SMA/WMA 使用)
|
* @param period 季节周期长度(SEASONAL 使用)
|
* @return 下一期预测
|
*/
|
public static Forecast forecast(String model, List<BigDecimal> history,
|
int window, int period) {
|
List<BigDecimal> values = CollUtil.isNotEmpty(history)
|
? history.stream().filter(v -> v != null).toList()
|
: List.of();
|
if (values.isEmpty()) {
|
return null;
|
}
|
String m = model == null ? "SMA" : model.toUpperCase();
|
BigDecimal value = switch (m) {
|
case "WMA" -> wma(values, window > 0 ? window : 12);
|
case "LR" -> linearRegressionForecast(values, 1);
|
case "SEASONAL" -> seasonalForecast(values, period > 0 ? period : 7);
|
case "YOY" -> yoyForecast(values);
|
default -> sma(values, window > 0 ? window : 12);
|
};
|
if (value == null) {
|
return null;
|
}
|
// 95% 置信区间:±10% 幅度
|
BigDecimal bound = value.abs().multiply(new BigDecimal("0.10"));
|
return new Forecast(value, value.subtract(bound), value.add(bound));
|
}
|
|
/**
|
* 简单移动平均
|
*/
|
public static BigDecimal sma(List<BigDecimal> values, int window) {
|
if (CollUtil.isEmpty(values)) {
|
return null;
|
}
|
int n = Math.min(window, values.size());
|
int from = values.size() - n;
|
BigDecimal sum = BigDecimal.ZERO;
|
for (int i = from; i < values.size(); i++) {
|
sum = sum.add(values.get(i));
|
}
|
return sum.divide(BigDecimal.valueOf(n), 4, RoundingMode.HALF_UP);
|
}
|
|
/**
|
* 加权移动平均:近期权重递增
|
*/
|
public static BigDecimal wma(List<BigDecimal> values, int window) {
|
if (CollUtil.isEmpty(values)) {
|
return null;
|
}
|
int n = Math.min(window, values.size());
|
int from = values.size() - n;
|
BigDecimal weightSum = BigDecimal.ZERO;
|
BigDecimal sum = BigDecimal.ZERO;
|
for (int i = from; i < values.size(); i++) {
|
int weight = i - from + 1;
|
sum = sum.add(values.get(i).multiply(BigDecimal.valueOf(weight)));
|
weightSum = weightSum.add(BigDecimal.valueOf(weight));
|
}
|
return sum.divide(weightSum, 4, RoundingMode.HALF_UP);
|
}
|
|
/**
|
* 一元线性回归(最小二乘),预测未来 extend 期
|
*/
|
public static BigDecimal linearRegressionForecast(List<BigDecimal> values, int extend) {
|
int n = values.size();
|
if (n < 2) {
|
return values.get(0);
|
}
|
BigDecimal sumX = BigDecimal.ZERO;
|
BigDecimal sumY = BigDecimal.ZERO;
|
BigDecimal sumXY = BigDecimal.ZERO;
|
BigDecimal sumXX = BigDecimal.ZERO;
|
for (int i = 0; i < n; i++) {
|
BigDecimal x = BigDecimal.valueOf(i);
|
BigDecimal y = values.get(i);
|
sumX = sumX.add(x);
|
sumY = sumY.add(y);
|
sumXY = sumXY.add(x.multiply(y));
|
sumXX = sumXX.add(x.multiply(x));
|
}
|
BigDecimal nbd = BigDecimal.valueOf(n);
|
BigDecimal slope = nbd.multiply(sumXY).subtract(sumX.multiply(sumY))
|
.divide(nbd.multiply(sumXX).subtract(sumX.multiply(sumX)), 6, RoundingMode.HALF_UP);
|
BigDecimal intercept = sumY.subtract(slope.multiply(sumX)).divide(nbd, 6, RoundingMode.HALF_UP);
|
BigDecimal xNext = BigDecimal.valueOf(n - 1 + extend);
|
return slope.multiply(xNext).add(intercept).setScale(4, RoundingMode.HALF_UP);
|
}
|
|
/**
|
* 季节指数预测:按周期归一化,用周期内同相位均值外推
|
*/
|
public static BigDecimal seasonalForecast(List<BigDecimal> values, int period) {
|
int n = values.size();
|
if (n < period) {
|
// 数据不足一个周期,退化为线性回归
|
return linearRegressionForecast(values, 1);
|
}
|
// 计算整体均值(去趋势参考)
|
BigDecimal avg = sma(values, n);
|
if (avg == null || avg.compareTo(BigDecimal.ZERO) == 0) {
|
avg = BigDecimal.ONE;
|
}
|
// 周期均值(每期基准)
|
BigDecimal periodAvg = BigDecimal.ZERO;
|
int baseCount = n / period;
|
for (int i = 0; i < n; i++) {
|
periodAvg = periodAvg.add(values.get(i));
|
}
|
periodAvg = periodAvg.divide(BigDecimal.valueOf(n), 6, RoundingMode.HALF_UP);
|
|
// 各相位均值
|
List<BigDecimal> phaseAvg = new ArrayList<>(period);
|
List<Integer> phaseCount = new ArrayList<>(period);
|
for (int i = 0; i < period; i++) {
|
phaseAvg.add(BigDecimal.ZERO);
|
phaseCount.add(0);
|
}
|
for (int i = 0; i < n; i++) {
|
int phase = i % period;
|
phaseAvg.set(phase, phaseAvg.get(phase).add(values.get(i)));
|
phaseCount.set(phase, phaseCount.get(phase) + 1);
|
}
|
BigDecimal xNext = BigDecimal.valueOf(n - 1 + 1); // 下一期
|
int nextPhase = n % period;
|
BigDecimal phaseBase = phaseCount.get(nextPhase) > 0
|
? phaseAvg.get(nextPhase).divide(BigDecimal.valueOf(phaseCount.get(nextPhase)), 6, RoundingMode.HALF_UP)
|
: periodAvg;
|
// 线性趋势调整
|
BigDecimal trend = linearRegressionForecast(values, 1);
|
BigDecimal trendAdjust = trend.multiply(avg).divide(periodAvg, 6, RoundingMode.HALF_UP);
|
return phaseBase.add(trendAdjust.subtract(periodAvg)).setScale(4, RoundingMode.HALF_UP);
|
}
|
|
/**
|
* 同比/环比外推:最后一个周期相对前一个周期的变化率外推
|
*/
|
public static BigDecimal yoyForecast(List<BigDecimal> values) {
|
int n = values.size();
|
if (n < 3) {
|
return values.get(n - 1);
|
}
|
BigDecimal last = values.get(n - 1);
|
BigDecimal prev = values.get(n - 2);
|
BigDecimal before = values.get(n - 3);
|
if (prev == null || prev.compareTo(BigDecimal.ZERO) == 0) {
|
return last;
|
}
|
// 前一周期变化率
|
BigDecimal growth = prev.subtract(before).divide(prev, 6, RoundingMode.HALF_UP);
|
return last.multiply(BigDecimal.ONE.add(growth)).setScale(4, RoundingMode.HALF_UP);
|
}
|
|
}
|