如何在Gaussian KDE等高线图中显示各区域点数?附额外问题
问题与解决方案
主问题:在KDE等高线区域内显示包含的点数
你当前用固定网格直方图统计点数的方式不符合需求,因为等高线是基于密度的不规则区域,而非固定网格。正确思路是:先获取contourf生成的所有等高线闭合区域,再统计每个区域内的原始点数量,最后将数量标注在区域中心。
具体实现步骤:
- 提取
contourf对象的所有闭合路径(每个路径对应一个等高线区域)。 - 对每个原始点,判断其所属的等高线区域。
- 计算区域中心坐标,将统计得到的点数标注在对应位置。
额外问题:等高线图与散点图x轴范围不一致的原因及解决
原因是你生成xcontour时用了np.linspace(0, np.max(x), len(x)),而散点图的x轴默认匹配原始数据的完整范围(0到1,因为np.random.rand生成0-1区间的数),等高线图的x轴被限制到np.max(x)(略小于1),即使设置sharex=True,也会因等高线数据范围限制导致显示不一致。
解决方法:将xcontour和ycontour的范围设置为原始数据的完整区间(0到1),或直接用ax.set_xlim统一两个子图的轴范围。
修改后的完整代码
import matplotlib.pyplot as plt import numpy as np from scipy.stats import gaussian_kde from matplotlib.path import Path # 生成随机数据 x = np.random.rand(100) y = np.random.rand(100) # 创建带共享x轴的两个子图 fig, ax = plt.subplots(2, sharex=True) # 生成覆盖完整数据范围的网格(解决轴范围不一致问题) x_min, x_max = 0, 1 y_min, y_max = 0, 1 grid_size = 100 xcontour = np.linspace(x_min, x_max, grid_size) ycontour = np.linspace(y_min, y_max, grid_size) xv, yv = np.meshgrid(xcontour, ycontour) positions = np.vstack([xv.ravel(), yv.ravel()]) # 计算高斯核密度估计 values = np.vstack([x, y]) kernel = gaussian_kde(values) f = np.reshape(kernel(positions).T, xv.shape) # 绘制等高线图 contour = ax[1].contourf(xv, yv, f, cmap='Blues') # 统计每个等高线区域内的点数 points = np.vstack([x, y]).T regions = contour.collections point_counts = [] for region in regions: # 获取当前区域的所有路径 paths = region.get_paths() count = 0 for path in paths: # 判断每个点是否在当前路径围成的区域内 inside = Path(path).contains_points(points) count += np.sum(inside) point_counts.append(count) # 在每个等高线区域中心标注点数 for i, region in enumerate(regions): if point_counts[i] == 0: continue # 获取区域的边界框,计算中心坐标 bbox = region.get_paths()[0].get_extents() cx = (bbox.x0 + bbox.x1) / 2 cy = (bbox.y0 + bbox.y1) / 2 ax[1].text(cx, cy, str(point_counts[i]), ha='center', va='center', color='white', fontweight='bold') # 添加色条并设置标签 cbar = plt.colorbar(contour, ax=ax[1], shrink=.6) cbar.ax.set_ylabel('Density') # 设置子图标题和绘制散点图 ax[1].set_title('Density of points with point counts') ax[0].scatter(x, y, s=1, marker=',') ax[0].set_title('Raw scatter points') # 统一轴范围(确保一致) ax[0].set_xlim(x_min, x_max) ax[1].set_xlim(x_min, x_max) plt.tight_layout() plt.show()
内容的提问来源于stack exchange,提问作者Luigi
相关产品推荐
相关产品推荐

