面向批量数据/颜色通道的在线方差更新高效算法求解
批量在线更新多维数据方差的高效实现方法
直接合并批次方差会导致结果偏差,正确的做法是通过跟踪总样本数、总体均值、总体平方均值来在线计算方差——因为方差的本质是 E[X²] - (E[X])²,只要维护这三个统计量,就能高效批量更新,完全适配多维数据(比如图像的颜色通道、像素维度),比单值的Welford算法快几个数量级。
实现步骤
- 初始化三个统计量:
n:已处理的总样本数mean:已处理数据的总体均值(维度和颜色通道一致)mean_sq:已处理数据的总体平方均值(维度和颜色通道一致)
- 对每个批次:
- 计算当前批次的样本数
batch_n、批次均值batch_mean、批次平方均值batch_mean_sq - 用加权平均的方式更新总均值和总平方均值,避免重新计算所有历史数据
- 计算当前批次的样本数
- 最终方差由
mean_sq - mean²得到(对应numpy默认的总体方差,若需样本方差则除以n-1)
修改后的示例代码
import numpy as np np.random.seed(0) # 正确的方差计算方式 all_data = np.random.randint(0, 9, (9, 3)) # <-- 无法存入内存 correct_var = all_data.var(axis=0) # 生成批次数据 batches = all_data.reshape(-1, 3, 3) # 初始化统计量 n = 0 mean = np.zeros(3) mean_sq = np.zeros(3) for batch in batches: batch_n = batch.shape[0] # 计算当前批次的均值和平方均值 batch_mean = batch.mean(axis=0) batch_mean_sq = (batch ** 2).mean(axis=0) # 更新总样本数 n += batch_n # 加权更新总体均值 mean = (mean * (n - batch_n) + batch_mean * batch_n) / n # 加权更新总体平方均值 mean_sq = (mean_sq * (n - batch_n) + batch_mean_sq * batch_n) / n # 计算在线方差 online_var = mean_sq - mean ** 2 # 验证结果一致 assert np.allclose(correct_var, online_var) print("结果一致,验证通过!")
关键说明
- 这个方法是全向量化的,所有操作都基于numpy的批量计算,完全适配多维数据(比如把示例中的(9,3)换成(10000, 256,256,3)的图像数据,代码无需修改)
- 内存占用仅取决于统计量的维度(颜色通道数),和总数据量无关,完美解决内存不足问题
- 结果和一次性计算的方差完全一致(浮点精度范围内)
内容的提问来源于stack exchange,提问作者Daraan
相关产品推荐
相关产品推荐

