如何在Python中按Bin内标签阈值占比为2D Histogram上色?
当然可以实现这个需求!其实不用完全手动显式遍历网格计算占比,我们有两种很方便的方案,适配不同的使用习惯,我给你详细拆解下:
方法1:用scipy.stats.binned_statistic_2d自定义统计函数
这种方法最贴合你之前的使用习惯,直接复用binned_statistic_2d的bin划分逻辑,只需要自定义一个统计函数来计算每个bin内标签低于阈值的占比即可。
首先定义一个统计函数,输入每个bin里的标签值,输出低于阈值的比例(注意处理空bin避免除以0):
def below_threshold_fraction(values, threshold=0.5): if len(values) == 0: return 0.0 # 空bin返回0,防止除以0报错 return (values < threshold).mean()
然后调用binned_statistic_2d时,把statistic参数设为这个自定义函数:
import numpy as np from scipy.stats import binned_statistic_2d import matplotlib.pyplot as plt # 替换成你的真实数据 x = np.random.randn(1000) y = np.random.randn(1000) labels = np.random.rand(1000) # 0-1的连续标签 # 设置bin数量或自定义边界(和你之前的设置保持一致) nbins = 20 threshold = 0.5 # 计算每个bin的占比 stat, xedges, yedges, _ = binned_statistic_2d( x, y, values=labels, statistic=below_threshold_fraction, bins=nbins, # 如果需要自定义bin边界,把bins换成[x_edges_list, y_edges_list]即可 ) # 可视化结果 plt.figure(figsize=(8,6)) plt.imshow(stat.T, origin='lower', extent=[xedges[0], xedges[-1], yedges[0], yedges[-1]], cmap='viridis') plt.colorbar(label=f'Fraction of points with label < {threshold}') plt.xlabel('X') plt.ylabel('Y') plt.title('2D Histogram Colored by Label Below Threshold Fraction') plt.show()
方法2:用numpy.histogram2d分别统计总点数和符合条件的点数
这种方法更直观,通过两次直方图统计,分别得到每个bin的总点数和标签低于阈值的点数,再相除得到占比,适合需要同时查看总点数分布的场景。
import numpy as np import matplotlib.pyplot as plt # 替换成你的真实数据 x = np.random.randn(1000) y = np.random.randn(1000) labels = np.random.rand(1000) threshold = 0.5 nbins = 20 # 统计每个bin的总点数 count_total, xedges, yedges = np.histogram2d(x, y, bins=nbins) # 统计每个bin内标签低于阈值的点数:只筛选符合条件的点做直方图 count_below, _, _ = np.histogram2d(x[labels < threshold], y[labels < threshold], bins=[xedges, yedges]) # 计算占比,处理空bin(避免除以0) fraction = np.divide(count_below, count_total, where=count_total != 0) fraction[count_total == 0] = 0.0 # 空bin直接设为0 # 可视化 plt.figure(figsize=(8,6)) plt.imshow(fraction.T, origin='lower', extent=[xedges[0], xedges[-1], yedges[0], yedges[-1]], cmap='viridis') plt.colorbar(label=f'Fraction of points with label < {threshold}') plt.xlabel('X') plt.ylabel('Y') plt.title('2D Histogram Colored by Label Below Threshold Fraction') plt.show()
两种方法的小提示
- 两种方法的bin设置要和你之前绘制标准差图时保持一致,这样结果可以直接对比。
- 空bin的处理很重要,两种方法都做了避免除以0的处理,你可以根据需求调整空bin的默认值(比如设为NaN,可视化时会显示空白)。
- 如果数据量特别大,
numpy.histogram2d的速度会略快一点;如果想和之前的代码逻辑完全统一,方法1会更顺手。
内容的提问来源于stack exchange,提问作者user8188120
相关产品推荐
相关产品推荐

