如何在Matplotlib中创建允许点重叠的类蜂群分布图
自定义类蜂群图的优化需求与问题
需求与问题描述
我想要实现一种类蜂群图(swarmplot)的图表,核心要求是:
- 清晰展示数据分布形态
- 支持通过点重叠方式快速绘制数万个数据点,理想效果如下:

我的实现思路是:将每个分布划分为分位数,对数据点的水平位置应用与该分位数内点数成比例的抖动。当各分布样本量相同时该方案可行,但当某一分布仅含少量数据点时,当前逻辑会导致点过度分散的错误效果,我需要调整抖动缩放逻辑,让少量数据点排列成近乎垂直的直线,错误效果如下:
当前实现代码
import matplotlib.pyplot as plt import numpy as np def fancy_distribution_plot(distributions: list, tick_labels: list, max_plot_width: int = 1, alpha=0.7, number_of_segments=12, separation_between_plots=0.1, separation_between_subplots=0.1, vertical_limits=None, grid=False, remove_outlier_above_segment=None, remove_outlier_below_segment=None, y_label=None, title=None): fig, ax = plt.subplots() number_of_plots = len(distributions) ax.set_xlim(left=0, right=number_of_plots * (max_plot_width + separation_between_plots) + separation_between_plots) ticks = [separation_between_plots + max_plot_width / 2 + (max_plot_width + separation_between_plots) * i for i in range(0, number_of_plots)] print(ticks) for i in range(len(distributions)): distribution = distributions[i] segments = np.linspace(np.min(distribution), np.max(distribution), number_of_segments + 1)[1:-1] segment_indices = number_of_segments - 1 - np.where(segments[:, None] >= distribution[None, :], 1, 0).sum(0) if remove_outlier_above_segment: a = remove_outlier_above_segment[i] distribution = distribution[segment_indices <= a] segment_indices = segment_indices[segment_indices <= a] if remove_outlier_below_segment: b = remove_outlier_below_segment[i] distribution = distribution[segment_indices >= b - 1] segment_indices = segment_indices[segment_indices >= b - 1] values, counts = np.unique(segment_indices, return_counts=True) counts_filled = [] j = 0 for k in range(number_of_segments): if k in values: counts_filled.append(counts[j]) j += 1 else: counts_filled.append(0) variances = (max_plot_width / 2) * (counts_filled / np.max(counts)) jitter_unadjusted = np.random.uniform(-1, 1, len(distribution)) jitter = np.take(variances, segment_indices) * jitter_unadjusted ax.scatter(jitter + ticks[i], distribution, alpha=alpha) ax.set_xticks(ticks) ax.set_xticklabels(tick_labels) if vertical_limits: ax.set_ylim(bottom=vertical_limits[0], top=vertical_limits[1]) if not grid: ax.grid(False) if y_label: ax.set_ylabel(y_label) if title: ax.set_title(title) plt.show()
错误效果复现代码
# 创建示例随机数据 np.random.seed(0) distro1 = np.random.normal(0, 2, 4) distr2 = np.random.normal(1, 1, 10) distr3 = np.random.normal(2, 3, 1000) distributions = [distro1, distr2, distr3] fancy_distribution_plot(distributions, tick_labels=['distro1', 'distro2', 'distro3'], number_of_segments=100, grid=False)
内容的提问来源于stack exchange,提问作者ufghd34
相关产品推荐
相关产品推荐

