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

基于现有双特征PDP热力图代码实现3特征4D PDP交互可视化咨询

4D PDP特征交互可视化实现说明

完全可以基于你现有的双特征PDP计算逻辑扩展实现需求效果:3个坐标轴对应三个输入特征,预测值均值(第4维度)同时通过曲面高度、表面热力颜色两种方式呈现。

实现核心修改点

  • 扩展原有双特征采样逻辑为三维网格采样,生成三个特征的等间隔采样点矩阵
  • 改用matplotlib.mplot3d模块绘制3D曲面,曲面高度直接映射预测值均值
  • 曲面表面启用颜色映射,用不同色值同步呈现预测值大小,实现双编码的4D呈现效果

可运行修改后代码

import numpy as np
import matplotlib.pyplot as plt
from matplotlib import cm
import logging

def pdp_custom_4D(estimated_model, X, y, n_splits, target_name_1, target_name_2, target_name_3, prefit=True, X_train=None, y_train=None):
    # 模型预训练校验
    if not prefit:
        try:
            estimated_model.fit(X_train, y_train)
        except:
            logging.warning("Estimated model must have fit method.")
    
    # 计算三个特征的取值范围
    f1_min, f1_max = X[target_name_1].min(), X[target_name_1].max()
    f2_min, f2_max = X[target_name_2].min(), X[target_name_2].max()
    f3_min, f3_max = X[target_name_3].min(), X[target_name_3].max()
    
    # 生成等间隔采样网格
    f1_grid = np.linspace(f1_min, f1_max, n_splits+1)
    f2_grid = np.linspace(f2_min, f2_max, n_splits+1)
    f3_grid = np.linspace(f3_min, f3_max, n_splits+1)
    X, Y, Z = np.meshgrid(f1_grid, f2_grid, f3_grid, indexing='ij')
    
    # 计算网格点对应PDP预测均值
    pdp_values = np.zeros_like(X)
    X_copy = X.copy()
    for i in range(n_splits+1):
        for j in range(n_splits+1):
            for k in range(n_splits+1):
                X_copy[target_name_1] = f1_grid[i]
                X_copy[target_name_2] = f2_grid[j]
                X_copy[target_name_3] = f3_grid[k]
                preds = estimated_model.predict(X_copy)
                pdp_values[i,j,k] = preds.mean()
    
    # 绘制3D曲面+热力映射
    fig = plt.figure(figsize=(12,8))
    ax = fig.add_subplot(111, projection='3d')
    # 取特征3的中位数截面绘制曲面(也可按需切换其他截面或添加滑动交互)
    mid_idx = n_splits//2
    surf = ax.plot_surface(X[:,:,mid_idx], Y[:,:,mid_idx], pdp_values[:,:,mid_idx], 
                           cmap=cm.coolwarm, edgecolor='none', alpha=0.8)
    # 坐标轴与图例设置
    ax.set_xlabel(target_name_1)
    ax.set_ylabel(target_name_2)
    ax.set_zlabel('Prediction Mean')
    fig.colorbar(surf, shrink=0.5, aspect=5, label='Prediction Value')
    plt.show()
    
    return

可选优化方向

  • 可添加交互滑块动态切换第三个特征的采样截面,实现完整三维特征空间的遍历查看
  • 调整cmap参数更换配色方案,或添加等高线投影增强数值辨识度
  • 若算力充足可将n_splits参数调大,提升曲面平滑度

内容的提问来源于stack exchange,提问作者PrisonedIbk9

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.07 08:00:03