如何使用plotly.express.imshow的facet_row参数实现分面热力图?
问题描述
尝试用plotly.express.imshow的facet_col参数绘制分面热力图,用Pandas生成了带多索引的DataFrame,但运行时出现AttributeError: 'DataFrame' object has no attribute 'dims'错误。用循环生成的4张热力图相互独立,轴没有关联,想知道能不能像px.scatter那样基于某列用facet_row/facet_col实现分面。
测试代码:
import pandas as pd import numpy as np import plotly.express as px # Create the index for the data frame x = np.linspace(-1,1, 6) y = np.linspace(-1,1,6) n_channel = [1, 2, 3, 4] xx, yy = np.meshgrid(x, y) zzz = np.random.randn(len(y)*len(n_channel),len(x)) df = pd.DataFrame( zzz, columns = pd.Index(x, name='x (m)'), index = pd.MultiIndex.from_product([y, n_channel], names=['y (m)', 'n_channel']), ) print(df) fig = px.imshow( df.reset_index('n_channel'), facet_col = 'n_channel', ) fig.write_html( 'plot.html', include_plotlyjs = 'cdn', )
循环实现代码:
for n in n_channel: fig = px.imshow( df.query(f'n_channel=={n}').reset_index('n_channel', drop=True), ) fig.write_html( f'plot_{n}.html', include_plotlyjs = 'cdn', )
解决方案
px.imshow的分面逻辑和px.scatter不同,它要求输入的数组/数据结构是三维张量(比如形状为(n_facets, y_dim, x_dim)的numpy数组),不能直接传入带额外分类列的DataFrame——这就是你报错的原因。
要实现轴关联的分面热力图,有两种可行方式:
方式一:用三维numpy数组直接生成分面(最简洁)
把不同n_channel的热力图数据堆叠成三维数组,直接传给px.imshow即可实现分面,且自动关联轴:
import pandas as pd import numpy as np import plotly.express as px # 生成数据 x = np.linspace(-1,1, 6) y = np.linspace(-1,1,6) n_channel = [1, 2, 3, 4] # 生成三维数据:维度为(分面数, y轴长度, x轴长度) zzz = np.random.randn(len(n_channel), len(y), len(x)) # 绘制分面热力图 fig = px.imshow( zzz, facet_col=0, # 按第0维度(n_channel)分面 x=x, y=y, labels=dict(facet_col="n_channel", x="x (m)", y="y (m)", color="value"), title="轴关联的分面热力图" ) # 统一所有子图的颜色范围,保证可视化一致性 fig.update_coloraxes(showscale=True, colorbar_x=-0.1) fig.write_html('plot_facet.html', include_plotlyjs='cdn')
方式二:用长格式DataFrame+Graph Objects手动构建(灵活性更高)
如果需要保留DataFrame结构,可以将数据转成长格式,用make_subplots创建子图,逐个添加热力图并关联轴:
import pandas as pd import numpy as np import plotly.graph_objects as go from plotly.subplots import make_subplots # 生成原始数据 x = np.linspace(-1,1, 6) y = np.linspace(-1,1,6) n_channel = [1, 2, 3, 4] xx, yy = np.meshgrid(x, y) zzz = np.random.randn(len(y)*len(n_channel),len(x)) df = pd.DataFrame( zzz, columns=pd.Index(x, name='x (m)'), index=pd.MultiIndex.from_product([y, n_channel], names=['y (m)', 'n_channel']), ).reset_index() # 创建共享轴的子图 fig = make_subplots( rows=1, cols=4, subplot_titles=[f"n_channel={n}" for n in n_channel], shared_xaxes=True, shared_yaxes=True, horizontal_spacing=0.05 ) # 获取全局颜色范围,保证所有子图颜色刻度一致 z_min = df.iloc[:, 2:].min().min() z_max = df.iloc[:, 2:].max().max() # 逐个添加热力图 for idx, n in enumerate(n_channel, 1): channel_data = df[df['n_channel'] == n].set_index('y (m)').drop('n_channel', axis=1) fig.add_trace( go.Heatmap( z=channel_data.values, x=channel_data.columns, y=channel_data.index, zmin=z_min, zmax=z_max, coloraxis="coloraxis" ), row=1, col=idx ) # 调整布局 fig.update_layout( coloraxis=dict(colorbar=dict(title="value")), title_text="轴关联的分面热力图", height=400, width=1000 ) fig.write_html('plot_facet_go.html', include_plotlyjs='cdn')
关键说明
px.imshow的分面依赖三维输入,这是它和px.scatter等统计图表的核心区别,不能直接传入带分类列的二维DataFrame。- 两种方式都能实现轴关联:第一种代码更简洁,适合快速实现;第二种可自定义子图布局、标题等细节,灵活性更强。
内容的提问来源于stack exchange,提问作者user171780
相关产品推荐
相关产品推荐

