使用SciPy binned_statistic计算分箱均值的NaN处理与报错解决
问题1:通用简单方案(兼容所有SciPy版本)
你之前代码输出结果不符合预期的核心原因是:过滤掉a中的nan后,x轴使用了重新生成的连续序列,和原始x轴的取值对应关系断裂,导致分箱区间和原始数据范围不匹配。
正确的通用做法是成对过滤x和待统计值的样本,仅删除待统计值为nan对应的(x, value)对,分箱区间仍采用原始x的取值范围,既不会触发SciPy的非有限值校验,也不会丢失你需要保留的分箱区间,自动过滤nan计算均值。
参考实现代码:
from scipy.stats import binned_statistic import numpy as np import matplotlib.pyplot as plt # 测试数据 a = np.array([0.1, 0.15, 0.17, 0.2, 0.3, 0.4, np.nan, 0.12, 0.15, 0.17, 0.22, np.nan, 0.37, np.nan, 0.12, 0.15, 0.17, 0.17, 0.35, 0.42, np.nan]) b = np.linspace(1, len(a), len(a)) # 成对过滤有效样本,保留原始x的对应取值 mask = ~np.isnan(a) valid_x = b[mask] valid_value = a[mask] # 分箱区间使用原始x的上下限,不会丢失需要保留的区间 bmean = binned_statistic(valid_x, valid_value, statistic='mean', bins=3, range=(b.min(), b.max())) # 计算分箱中心 bin_edges = bmean.bin_edges bin_centers = (bin_edges[1:] + bin_edges[:-1]) / 2 # 结果验证 plt.scatter(b, a, label='原始数据') plt.hlines(np.nanmean(a), b[0], b[-1], linestyles='--', label='全局均值') plt.scatter(bin_centers, bmean.statistic, marker='x', s=90, c='red', label='分箱均值') plt.legend() plt.show()
如果某个分箱内全为nan,返回的统计值为nan,你可以后续按需补充空分箱的填充逻辑。
问题2:直接使用statistic=np.nanmean的方法
根据你使用的SciPy版本,对应两种处理方式:
- SciPy >= 1.9.0版本:官方已修复历史校验逻辑问题,新增
nan_policy参数,设置为omit即可直接传入np.nanmean,不需要提前过滤数据:
# 直接传入原始含nan数据即可 bmean = binned_statistic(b, a, statistic=np.nanmean, bins=3, range=(b.min(), b.max()), nan_policy='omit')
- SciPy 1.4.0 ~ 1.8.x版本:这几个版本的输入校验强制要求所有传入值为有限值,无官方绕过参数,建议优先使用上面的通用过滤方案,或升级SciPy到1.9以上版本即可支持
nan_policy参数。
内容的提问来源于stack exchange,提问作者nuwe
相关产品推荐
相关产品推荐

