如何使用ojalgo最优计算滚动均值?现有实现存疑问
用ojAlgo高效计算滚动均值的最佳方案
我懂你的困扰——用SampleSet来算滚动均值确实会有重复计算的问题,尤其是窗口大的时候,每个元素被反复加好几次,效率实在拉胯。至于你说的第二种方案意图不明确,我来给你梳理下ojAlgo里计算滚动均值的最佳姿势!
为什么SampleSet的基础方案效率低?
你提到的第一种用SampleSet的方案,每次滑动窗口都重新创建一个包含当前窗口元素的SampleSet,再计算均值。这种方式的问题在于:
- 每个元素会在它所在的窗口里被重复累加,窗口越大,重复计算的次数越多,时间复杂度是O(n*k)(n是数据长度,k是窗口大小),数据量大的时候性能会明显下降;
- SampleSet本身是为统计整体样本特征设计的,不是专门针对滑动窗口场景优化的。
最佳方案1:前缀和法(仅需滚动均值时首选)
如果你的需求只是计算滚动均值,前缀和法是最高效的选择,时间复杂度只有O(n),核心思路是先计算前缀和数组,再通过两个前缀和的差值快速得到任意窗口的和,最后除以窗口大小得到均值。
代码实现
import org.ojalgo.array.Array1D; import org.ojalgo.array.Primitive32Array; import org.ojalgo.random.Uniform; public class RollingMeanWithPrefixSum { public static void main(String[] args) { // 生成测试数据(和你提供的代码一致) final Array1D<Double> doubles = Array1D.factory(Primitive32Array.FACTORY) .makeFilled(10, new Uniform()); System.out.println("原始数据: " + doubles.toString()); int windowSize = 3; // 可根据需求调整窗口大小 int dataSize = doubles.size(); // 计算前缀和数组:prefixSum[0]=0,prefixSum[i] = doubles[0]+...+doubles[i-1] Array1D<Double> prefixSum = Array1D.PRIMITIVE32.makeZero(dataSize + 1); for (int i = 0; i < dataSize; i++) { prefixSum.set(i + 1, prefixSum.get(i) + doubles.get(i)); } // 计算滚动均值 Array1D<Double> rollingMean = Array1D.PRIMITIVE32.makeZero(dataSize); for (int i = 0; i < dataSize; i++) { if (i < windowSize - 1) { // 前windowSize-1个元素窗口不完整,这里保留原始值,你也可以设为NaN或跳过 rollingMean.set(i, doubles.get(i)); } else { // 窗口完整时,用前缀和差值计算窗口和,再求均值 double windowSum = prefixSum.get(i + 1) - prefixSum.get(i - windowSize + 1); rollingMean.set(i, windowSum / windowSize); } } System.out.printf("滚动均值(窗口大小=%d): %s\n", windowSize, rollingMean.toString()); } }
代码细节说明
- 前缀和数组的设计是为了快速计算任意区间的和,避免重复累加;
- 前windowSize-1个元素的处理可根据业务需求调整:比如要求窗口必须完整,就从
i = windowSize - 1开始计算,前面的位置设为Double.NaN。
最佳方案2:优化后的SampleSet方案(需要多统计量时用)
如果你除了滚动均值,还需要窗口内的其他统计量(比如方差、中位数等),可以用优化后的SampleSet方案——每次滑动窗口时,只移除窗口最左侧的元素,加入新的右侧元素,而不是重新创建整个SampleSet,这样时间复杂度也能降到O(n)。
代码实现
import org.ojalgo.array.Array1D; import org.ojalgo.array.Primitive32Array; import org.ojalgo.random.Uniform; import org.ojalgo.statistics.SampleSet; public class OptimizedRollingMeanWithSampleSet { public static void main(String[] args) { Array1D<Double> doubles = Array1D.factory(Primitive32Array.FACTORY) .makeFilled(10, new Uniform()); System.out.println("原始数据: " + doubles.toString()); int windowSize = 3; int dataSize = doubles.size(); Array1D<Double> rollingMean = Array1D.PRIMITIVE32.makeZero(dataSize); // 初始化第一个完整窗口的SampleSet SampleSet sampleSet = new SampleSet(doubles.subList(0, windowSize)); rollingMean.set(windowSize - 1, sampleSet.getMean()); // 滑动窗口:移除最左元素,加入新元素 for (int i = windowSize; i < dataSize; i++) { sampleSet.remove(doubles.get(i - windowSize)); sampleSet.add(doubles.get(i)); rollingMean.set(i, sampleSet.getMean()); } // 处理前windowSize-1个不完整窗口的元素 for (int i = 0; i < windowSize - 1; i++) { rollingMean.set(i, doubles.get(i)); } System.out.printf("滚动均值(窗口大小=%d): %s\n", windowSize, rollingMean.toString()); } }
优势说明
- 复用同一个SampleSet实例,避免重复创建对象和遍历元素;
- 可以直接调用SampleSet的其他方法(比如
getVariance()、getMedian())获取窗口内的其他统计值,无需额外实现。
内容的提问来源于stack exchange,提问作者user482745
相关产品推荐
相关产品推荐

