如何创建单元格内嵌sparkline的分类热力图?
实现单元格内嵌Sparkline的分类热力图
下面是用Matplotlib实现每个单元格包含独立sparkline的热力图方案,核心思路是在基础热力图的每个单元格内嵌入小型子图绘制sparkline,替代传统文本标注:
完整代码示例
import matplotlib.pyplot as plt from mpl_toolkits.axes_grid1.inset_locator import inset_axes import numpy as np # 1. 准备模拟数据 row_labels = ["类别A", "类别B", "类别C", "类别D"] col_labels = ["组1", "组2", "组3", "组4"] n_rows = len(row_labels) n_cols = len(col_labels) # 热力图主数据(控制单元格颜色) heatmap_data = np.random.rand(n_rows, n_cols) * 10 - 5 # 每个单元格对应的sparkline时间序列数据(每个单元格生成10个点的序列) sparkline_data = [ [np.random.randn(10).cumsum() for _ in range(n_cols)] for _ in range(n_rows) ] # 2. 绘制基础热力图 fig, ax = plt.subplots(figsize=(10, 6)) im = ax.imshow(heatmap_data, cmap="coolwarm", aspect="equal") # 设置热力图标签与刻度 ax.set_xticks(np.arange(n_cols)) ax.set_yticks(np.arange(n_rows)) ax.set_xticklabels(col_labels) ax.set_yticklabels(row_labels) ax.tick_params(top=True, bottom=False, labeltop=True, labelbottom=False) # 添加颜色条 plt.colorbar(im, ax=ax, shrink=0.8) # 3. 在每个单元格内嵌Sparkline for i in range(n_rows): for j in range(n_cols): # 创建内嵌子图,占单元格80%的空间 inset_ax = inset_axes( ax, width="80%", height="80%", loc="center", bbox_to_anchor=(j, i, 1, 1), bbox_transform=ax.transData, borderpad=0 ) # 绘制sparkline data = sparkline_data[i][j] inset_ax.plot(data, color="black", linewidth=1.5) # 隐藏轴元素,简化样式 inset_ax.set_xticks([]) inset_ax.set_yticks([]) for spine in inset_ax.spines.values(): spine.set_visible(False) # 可选:标记最大值点 max_idx = np.argmax(data) inset_ax.scatter(max_idx, data[max_idx], color="red", s=15) # 4. 调整整体布局 plt.tight_layout() plt.show()
关键实现技巧
- 内嵌子图定位:用
inset_axes通过bbox_to_anchor和transData精准绑定单元格位置,确保sparkline与单元格对齐。 - Sparkline简化:隐藏坐标轴、边框,只保留核心线条,避免干扰热力图视觉;可按需添加极值点、趋势标记提升信息密度。
- 热力图基础优化:设置
aspect="equal"保证单元格为正方形,调整刻度位置让分类标签显示更合理。
简化版方案(基于Annotation)
如果不想用内嵌子图,也可以用PathPatch结合annotate绘制极简sparkline,灵活性稍弱但代码更轻量:
from matplotlib.path import Path from matplotlib.patches import PathPatch # 在单元格(j,i)绘制简化sparkline data = sparkline_data[i][j] # 归一化数据到单元格内部范围 norm_data = (data - data.min()) / (data.max() - data.min()) * 0.8 + 0.1 vertices = [(j + x/9, i + y) for x, y in enumerate(norm_data)] path = Path(vertices) patch = PathPatch(path, color="black", linewidth=1, fill=False) ax.add_patch(patch)
内容的提问来源于stack exchange,提问作者As3adTintin
相关产品推荐
相关产品推荐

