You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

基于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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.19 10:32:06