如何提升数组最大值、均值、标准差等基础统计量的获取效率
数组基础统计量获取效率优化方案
1 核心复杂度优化:将O(nlogn)降为O(n)
你当前的性能瓶颈完全来自于为了求中位数做的全量排序,实际上求中位数不需要全排序,用*快速选择(Quickselect)*算法即可:
- 平均时间复杂度为O(n),最坏情况可以配合中位数中位数(Median of medians)算法稳定在O(n),远优于全排序的O(nlogn)
- 注意不要直接修改原输入数组,需要先复制一份再做快速选择,避免污染原有数据(你现有代码的
Arrays.sort是原地排序,会改动传入的数组,属于隐性bug)
2 减少遍历次数:一次遍历拿到4个指标
原来的均值、标准差计算走了两次遍历,还单独找max、min,完全可以合并成一次遍历:
- 用Welford在线方差算法,一次遍历就能算出均值和样本标准差,数值稳定性也比你现在先算均值再算平方差的方案更高
- 遍历过程中顺便记录max和min,不需要额外遍历
- 遍历完成后只需要额外做一次快速选择求中位数,整体只需要2次线性遍历,远优于原有方案的3次+全排序
3 工程层面优化(适配批量处理场景)
- 去掉多余的中间数组创建:原有代码每次计算都生成2个临时double数组,批量场景下会增加GC开销,直接在统计方法内赋值结果即可
- 空数组、长度为1的边界判断统一放在入口方法,避免重复判断
- 如果数组长度固定,可以提前预分配临时数组复用,不需要每次计算都新建
- 数值计算替换
Math.pow为直接乘法,Math.pow的开销远高于(num - mean)*(num - mean)
优化后代码示例
import java.time.Duration; import java.time.Instant; import java.util.Arrays; import java.util.Random; public class StatsOptimization { // 快速选择找第k小的元素 private static double quickSelect(double[] arr, int k) { // 复制数组避免修改原数组 double[] copy = Arrays.copyOf(arr, arr.length); int left = 0, right = copy.length - 1; Random rand = new Random(); while (left < right) { int pivotIdx = left + rand.nextInt(right - left); pivotIdx = partition(copy, left, right, pivotIdx); if (k < pivotIdx) { right = pivotIdx - 1; } else if (k > pivotIdx) { left = pivotIdx + 1; } else { break; } } return copy[k]; } private static int partition(double[] arr, int left, int right, int pivotIdx) { double pivotVal = arr[pivotIdx]; swap(arr, pivotIdx, right); int storeIdx = left; for (int i = left; i < right; i++) { if (arr[i] < pivotVal) { swap(arr, storeIdx, i); storeIdx++; } } swap(arr, right, storeIdx); return storeIdx; } private static void swap(double[] arr, int i, int j) { double temp = arr[i]; arr[i] = arr[j]; arr[j] = temp; } public static double[] getSummaryStatistics(double[] a) { int len = a.length; if (len == 0) { throw new IllegalArgumentException("Array is empty, please verify the values."); } double[] summary = new double[5]; if (len == 1) { Arrays.fill(summary, a[0]); summary[2] = 0; return summary; } // 一次遍历算max、min、均值、样本标准差 double max = -Double.MAX_VALUE; double min = Double.MAX_VALUE; double mean = 0; double m2 = 0; // Welford算法中间变量 for (int i = 0; i < len; i++) { double val = a[i]; // 更新max min if (val > max) max = val; if (val < min) min = val; // Welford更新 double delta = val - mean; mean += delta / (i + 1); double delta2 = val - mean; m2 += delta * delta2; } double std = Math.sqrt(m2 / (len - 1)); // 快速选择算中位数 double median; if (len % 2 == 1) { median = quickSelect(a, len / 2); } else { median = (quickSelect(a, len/2 - 1) + quickSelect(a, len/2)) / 2.0; } // 赋值结果:0=max,1=mean,2=std,3=median,4=min summary[0] = max; summary[1] = mean; summary[2] = std; summary[3] = median; summary[4] = min; return summary; } // 测试方法 public static void main(String[] args) { int numVals=1000; double[] ar=new double[numVals]; int numCalculations=2*1000*1*1000; Instant start = Instant.now(); Random random = new Random(); for(int i=0;i<numCalculations;i++){ for(int k=0;k<numVals;k++){ ar[k]= random.nextDouble(); } double[] stats=getSummaryStatistics(ar); } Instant end = Instant.now(); long totalTime = Duration.between(start, end).toSeconds(); System.out.println("Time (s):" + totalTime); } }
额外优化建议
如果允许近似结果,对精度要求不高的场景,可以用近似中位数算法,中位数计算速度可以再提升3-5倍,误差可以控制在1%以内,非常适合超大规模批量统计场景。
内容的提问来源于stack exchange,提问作者DTK
相关产品推荐
相关产品推荐

