3 天以前 5e0640513226d9d9f2d766c075f79832c9d290ba
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
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);
    }
 
}