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

Statsmodels二维Nadaraya-Watson核回归实现求助

解决Statsmodels中二维Nadaraya-Watson核回归的问题

我来帮你搞定这个二维Nadaraya-Watson回归的问题!你之前做一维回归没问题,二维卡壳的核心原因是自变量的输入格式和var_type参数的设置,下面给你完整的可运行代码和关键细节解释:

完整代码示例

首先确保你导入了必要的库,然后按以下步骤操作:

import numpy as np
import pandas as pd
import statsmodels.nonparametric.kernel_regression as kr

# 假设你的目标DataFrame是XYZ,包含X、Y、Z三列
# 1. 准备自变量和因变量:自变量要合并成二维数组
exog = XYZ[['X', 'Y']].values  # 形状为(n_samples, 2)的二维数组
endog = XYZ['Z'].values        # 因变量是一维数组

# 2. 初始化KernelReg模型
# var_type='cc'表示两个自变量都是连续变量(continuous),如果有分类变量用'd'
model = kr.KernelReg(endog=endog, exog=exog, var_type='cc')

# 3. 拟合模型并获取预测结果
# 对原始数据预测,得到拟合值和标准差
fitted_vals, std_errs = model.fit()

# (可选)生成网格数据可视化二维回归面
# 创建X和Y的网格
x_grid = np.linspace(XYZ['X'].min(), XYZ['X'].max(), 100)
y_grid = np.linspace(XYZ['Y'].min(), XYZ['Y'].max(), 100)
x_mesh, y_mesh = np.meshgrid(x_grid, y_grid)

# 把网格数据转换成模型需要的二维输入格式
grid_input = np.column_stack((x_mesh.ravel(), y_mesh.ravel()))
# 预测网格点的Z值
grid_preds, _ = model.fit(grid_input)
# 把预测值转换成网格形状,方便绘图
grid_preds = grid_preds.reshape(x_mesh.shape)

关键细节解释

  • 自变量格式:一维回归时你用的是单变量的一维数组,但二维回归必须把X和Y合并成形状为(n_samples, 2)的二维数组,这是最容易踩坑的地方!
  • var_type参数:这个参数要和自变量的数量、类型一一对应。两个连续变量就写'cc',如果其中一个是分类变量就改成'cd'或'dc',字母顺序要和自变量的顺序一致。
  • 模型配置:你可以手动指定核函数(比如kernel='epanechnikov',默认是高斯核),或者调整带宽(默认是自动交叉验证选择最优带宽,你可以用model.bw查看自动选择的带宽值,也可以通过bw参数手动设置)。

验证结果(可选)

你可以把拟合值和原始Z值对比,或者用matplotlib绘制三维回归面来直观查看X和Y对Z分布的影响:

import matplotlib.pyplot as plt
from mpl_toolkits.mplot3d import Axes3D

fig = plt.figure(figsize=(10,7))
ax = fig.add_subplot(111, projection='3d')
# 绘制原始数据点
ax.scatter(XYZ['X'], XYZ['Y'], XYZ['Z'], alpha=0.5, label='Raw Data')
# 绘制回归面
ax.plot_surface(x_mesh, y_mesh, grid_preds, cmap='viridis', alpha=0.7, label='Regression Surface')
ax.set_xlabel('X')
ax.set_ylabel('Y')
ax.set_zlabel('Z')
plt.legend()
plt.show()

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.20 10:27:14