如何创建数量可变的Matplotlib子图?ML可视化代码优化
动态生成ML模型预测子图的优化方案
核心思路
把重复的子图配置逻辑封装成循环,通过模型配置列表统一管理所有模型的参数,新增模型时只需在列表中添加条目,无需修改大量重复代码。
优化后代码
import matplotlib.pyplot as plt # 1. 定义模型配置:(数据列名, 子图标签) model_configs = [ ("actual", "X"), ("rfr", "X_rfr"), ("gbr", "X_gbr"), ("knr", "X_knr"), ("lir", "X_lir"), ("rlr", "X_rlr"), ("llr", "X_llr"), ("enr", "X_enr"), ("svr", "X_svr"), ("krr", "X_krr"), ("brr", "X_brr"), ("par", "X_par"), ("gpr", "X_gpr"), ("sgd", "X_sgd"), ("mlp", "X_mlp") ] # 2. 动态创建子图:根据配置数量生成对应列数的子图 n_models = len(model_configs) fig, axes = plt.subplots(nrows=1, ncols=n_models, sharex=True, sharey=True, figsize=(15,15)) # 3. 循环配置每个子图 for idx, (col_name, xlabel) in enumerate(model_configs): ax = axes[idx] # 添加顶部边框(隐藏额外x轴) axt = ax.twiny() axt.xaxis.set_visible(False) # 基础样式配置 ax.grid(which='major', color='lightgrey', linestyle='-') ax.set_xlim(-3, 3) ax.set_xlabel(xlabel) ax.spines["top"].set_position(("axes", 1.02)) ax.invert_yaxis() # 绘图:第一个子图用绿色标记 if idx == 0: ax.plot(col_name, 'D', data=preds, color='green') else: ax.plot(col_name, 'D', data=preds) # 隐藏非第一个子图的y轴标签 plt.setp(ax.get_yticklabels(), visible=False) # 设置x轴刻度和标签在顶部 ax.xaxis.set_ticks_position("top") ax.xaxis.set_label_position("top") plt.tight_layout() plt.show()
关键优势
- 扩展性强:新增模型时,只需在
model_configs列表中添加(列名, 标签)元组即可 - 代码简洁:消除大量重复的子图配置代码,维护成本大幅降低
- 一致性高:所有子图样式统一,避免手动配置导致的差异
内容的提问来源于stack exchange,提问作者Mahmoud Shihab
相关产品推荐
相关产品推荐

