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 history, int window, int period) { List 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 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 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 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 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 phaseAvg = new ArrayList<>(period); List 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 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); } }