基于lmfit/scipy brute方法的多参数优化结果网格热图绘制优化
嘿,这个需求我刚好折腾过!针对多参数优化后的两两配对热图网格,我给你整理了一套可行的实现方案,结合lmfit的结果来做很顺畅。
实现思路与步骤
核心逻辑是:从lmfit暴力搜索结果中提取参数组合与目标函数值 → 生成所有无重复参数对 → 为每一对绘制热图(聚焦参数交互下的目标函数变化) → 拼成网格布局展示。
1. 先提取核心数据
假设你已经用lmfit完成了暴力搜索(不管是用BruteMinimizer还是结合scipy的brute方法),结果存在result对象里,先把关键数据提出来:
# 获取参数名称列表 param_names = list(result.params.keys()) # 获取所有参数组合的取值(每行是一组参数值) param_grids = result.brute_grid # 获取对应每组参数的目标函数值 objective_vals = result.brute_Jout
2. 生成无重复参数配对
用itertools.combinations一键生成所有两两不重复的参数对,完全符合你要的b~c、b~d这类组合:
import itertools # 生成所有2个参数的无重复配对 param_pairs = list(itertools.combinations(param_names, 2))
3. 绘制网格热图
这里用matplotlib+seaborn来实现,代码里我做了两种固定其他参数的方式(选适合你的就行):
完整代码示例
import numpy as np import matplotlib.pyplot as plt import seaborn as sns # 计算网格布局:自动适配配对数量,避免浪费空间 n_pairs = len(param_pairs) n_cols = int(np.ceil(np.sqrt(n_pairs))) n_rows = int(np.ceil(n_pairs / n_cols)) # 创建画布 fig, axes = plt.subplots(n_rows, n_cols, figsize=(4*n_cols, 4*n_rows)) axes = axes.flatten() # 把二维轴数组转成一维,方便遍历 # 遍历每个参数对绘制热图 for idx, (p1, p2) in enumerate(param_pairs): ax = axes[idx] # 获取两个参数在列表中的索引 idx1 = param_names.index(p1) idx2 = param_names.index(p2) # -------------------------- # 方式1:固定其他参数在最优值(推荐,聚焦最优解附近的参数交互) # -------------------------- other_params = [name for name in param_names if name not in (p1, p2)] mask = np.ones(len(param_grids), dtype=bool) # 筛选出其他参数等于最优值的组合 for name in other_params: param_idx = param_names.index(name) mask &= (param_grids[:, param_idx] == result.params[name].value) # -------------------------- # 方式2:对其他参数的所有组合取均值(展示全局参数交互) # -------------------------- # mask = np.ones(len(param_grids), dtype=bool) # 不需要筛选,用所有数据 # 提取当前参数对的取值和对应目标函数值 p1_vals = np.unique(param_grids[mask, idx1]) p2_vals = np.unique(param_grids[mask, idx2]) J_grid = np.zeros((len(p1_vals), len(p2_vals))) # 填充热图数据 for i, val1 in enumerate(p1_vals): for j, val2 in enumerate(p2_vals): pair_mask = (param_grids[mask, idx1] == val1) & (param_grids[mask, idx2] == val2) J_grid[i, j] = np.mean(objective_vals[mask][pair_mask]) # 绘制热图 sns.heatmap(J_grid, ax=ax, xticklabels=np.round(p2_vals, 3), yticklabels=np.round(p1_vals, 3), cmap='viridis', annot=False, cbar=True) ax.set_xlabel(p2) ax.set_ylabel(p1) ax.set_title(f'Objective: {p1} vs {p2}') # 隐藏多余的空白子图 for ax in axes[n_pairs:]: ax.axis('off') plt.tight_layout() plt.show()
4. 关键细节调整
- 热图优化:如果目标函数值差异大,可以加
norm=plt.LogNorm()做对数缩放;参数采样点少的话,开annot=True显示具体数值更直观。 - 暴力搜索分辨率:用lmfit的
BruteMinimizer时,记得设置Ns参数(比如Ns=20),每个参数采样足够多的点,热图才会清晰。 - 目标函数适配:如果你的优化是多目标,只需要把
objective_vals换成你关注的那个目标的结果即可。
示例效果
比如你有4个参数b、c、d、e,运行后会生成6个热图,自动排列成3x2的网格,每个子图对应一组参数对的目标函数分布,完全满足你要的可视化需求。
内容的提问来源于stack exchange,提问作者FriskyGrub
相关产品推荐
相关产品推荐

