Python实现:为热图类别列设专属底色并按数值调深浅
需求:带列专属底色的热图绘制
我需要制作一款热图,要求:
- X轴的5个类别各有专属的整体底色
- 每个月份对应的数值决定该单元格颜色的深浅
目前使用sns.heatmap绘制,现有代码如下,通过cmap="Greys"能呈现数值差异,但希望实现每列不同预设底色且不添加额外单元格:
sns.heatmap(viewed, linewidth=.5, cmap="Greys", fmt="d")
解决方案
方法一:底色叠加灰度数值层
通过先绘制列专属底色,再叠加基于数值的灰度层,实现底色固定、数值控制深浅的效果:
import seaborn as sns import matplotlib.pyplot as plt import numpy as np # 替换为你的实际数据矩阵,形状为(月份数, 类别数) viewed = np.random.randint(10, 100, (6, 5)) # 定义5个类别对应的专属底色 base_colors = ["#FF9999", "#99FF99", "#9999FF", "#FFCC99", "#CC99FF"] fig, ax = plt.subplots(figsize=(8, 4)) # 绘制列专属底色 for col in range(viewed.shape[1]): # 生成对应列的全1矩阵,用于填充底色 col_base = np.ones((viewed.shape[0], 1)) ax.imshow(col_base, extent=[col, col+1, viewed.shape[0], 0], cmap=plt.cm.colors.ListedColormap([base_colors[col]]), alpha=1) # 归一化数值到0-1区间,用于控制灰度层透明度 norm = plt.Normalize(viewed.min(), viewed.max()) normed_data = norm(viewed) # 反转灰度逻辑:数值越大,灰度越浅,底色显示越清晰 gray_data = 1 - normed_data # 叠加数值灰度层 ax.imshow(gray_data, cmap="Greys", extent=[0, viewed.shape[1], viewed.shape[0], 0], alpha=gray_data, interpolation="nearest") # 设置刻度与边框 ax.set_xticks(np.arange(viewed.shape[1]) + 0.5) ax.set_xticklabels(["类别1", "类别2", "类别3", "类别4", "类别5"]) ax.set_yticks(np.arange(viewed.shape[0]) + 0.5) ax.set_yticklabels(["1月", "2月", "3月", "4月", "5月", "6月"]) ax.grid(color="white", linewidth=0.5) plt.tight_layout() plt.show()
方法二:列专属渐变色彩映射
为每个类别创建独立的渐变色彩映射,数值直接控制该底色的深浅程度:
import seaborn as sns import matplotlib.pyplot as plt from matplotlib.colors import LinearSegmentedColormap import numpy as np # 替换为你的实际数据 viewed = np.random.randint(10, 100, (6, 5)) # 定义5个类别的基础底色 base_colors = ["#FF9999", "#99FF99", "#9999FF", "#FFCC99", "#CC99FF"] fig, ax = plt.subplots(figsize=(8, 4)) # 为每个类别生成专属渐变色彩映射(浅到深) custom_cmaps = [] for color in base_colors: # 生成基础色的浅色调(低数值)和深色调(高数值) light_shade = plt.cm.colors.to_rgba(color, alpha=0.3) dark_shade = plt.cm.colors.to_rgba(color, alpha=1.0) cmap = LinearSegmentedColormap.from_list(f"col_cmap_{color}", [light_shade, dark_shade]) custom_cmaps.append(cmap) # 逐列绘制热图 for col_idx in range(viewed.shape[1]): col_data = viewed[:, col_idx].reshape(-1, 1) # 绘制当前列,使用专属色卡 sns.heatmap(col_data, ax=ax, cmap=custom_cmaps[col_idx], fmt="d", cbar=False, linewidths=0.5, xticklabels=[f"类别{col_idx+1}"], # 仅第一列显示Y轴刻度,避免重复 yticklabels=["1月", "2月", "3月", "4月", "5月", "6月"] if col_idx == 0 else False, extent=[col_idx, col_idx+1, viewed.shape[0], 0]) plt.tight_layout() plt.show()
内容的提问来源于stack exchange,提问作者Matthew JJJ
相关产品推荐
相关产品推荐

