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

kNN分类器可视化:原始特征值显示及坐标轴刻度调整问题

解决方案:用原始特征空间可视化kNN决策边界+悬停显示数值

嗨,这个需求我之前做项目时也碰到过,手动改刻度确实能凑合用,但其实有更准确流畅的方案,同时悬停显示原始数值的问题也能轻松解决,我给你一步步拆解:

一、更优的可视化方案:直接在原始特征空间绘制决策边界

你之前想改xticks/yticks的思路没问题,但容易出现决策边界位置和原始特征坐标不匹配的问题。更好的做法是:用标准化特征训练模型,但把用于计算决策边界的网格点逆变换回原始特征空间,然后直接在原始空间绘制所有元素,这样坐标轴天然就是原始的年龄和薪资值,完全不需要手动调整刻度。

举个代码示例(假设用StandardScaler标准化):

import numpy as np
import matplotlib.pyplot as plt
from sklearn.neighbors import KNeighborsClassifier
from sklearn.preprocessing import StandardScaler

# 模拟原始数据(年龄、薪资)和标签
X_original = np.array([[22, 42000], [27, 55000], [33, 68000], [38, 75000], [45, 92000], [50, 105000]])
y = np.array([0, 0, 1, 1, 1, 0])

# 标准化特征用于训练模型
scaler = StandardScaler()
X_scaled = scaler.fit_transform(X_original)

# 训练kNN分类器
knn = KNeighborsClassifier(n_neighbors=3)
knn.fit(X_scaled, y)

# 1. 创建标准化后的网格(用于计算决策边界)
h = 0.02  # 网格步长,越小边界越平滑
x_scaled_min, x_scaled_max = X_scaled[:, 0].min() - 0.5, X_scaled[:, 0].max() + 0.5
y_scaled_min, y_scaled_max = X_scaled[:, 1].min() - 0.5, X_scaled[:, 1].max() + 0.5
xx_scaled, yy_scaled = np.meshgrid(np.arange(x_scaled_min, x_scaled_max, h),
                                   np.arange(y_scaled_min, y_scaled_max, h))

# 2. 预测网格点的分类结果
Z = knn.predict(np.c_[xx_scaled.ravel(), yy_scaled.ravel()])
Z = Z.reshape(xx_scaled.shape)

# 3. 将标准化网格逆变换回原始特征空间
grid_original = scaler.inverse_transform(np.c_[xx_scaled.ravel(), yy_scaled.ravel()])
xx_original = grid_original[:, 0].reshape(xx_scaled.shape)
yy_original = grid_original[:, 1].reshape(xx_scaled.shape)

# 4. 在原始特征空间绘制决策边界和数据点
plt.figure(figsize=(8, 6))
# 绘制决策边界
plt.contourf(xx_original, yy_original, Z, alpha=0.3, cmap=plt.cm.Paired)
# 绘制原始数据点
plt.scatter(X_original[:, 0], X_original[:, 1], c=y, edgecolors='k', s=100, cmap=plt.cm.Paired)

# 设置坐标轴标签
plt.xlabel('*年龄*', fontsize=12)
plt.ylabel('*薪资*', fontsize=12)
plt.title('kNN分类器决策边界(原始特征空间)', fontsize=14)

这样出来的图表,所有元素都是基于原始年龄和薪资的坐标,完全符合你的需求,而且不会出现刻度和边界错位的问题。

二、实现鼠标悬停显示原始数值

要实现鼠标移动时显示当前位置的年龄和薪资,可以利用Matplotlib的事件处理机制,创建一个动态更新的注释框:

在上面的代码末尾添加以下内容:

# 创建隐藏的注释框,用于显示悬停信息
annot = plt.annotate(
    "", xy=(0,0), xytext=(-100, 20),  # xytext设置注释框位置(左下角附近)
    textcoords="offset points",
    bbox=dict(boxstyle="round", fc="white", alpha=0.9),
    arrowprops=dict(arrowstyle="->", alpha=0.7)
)
annot.set_visible(False)

def update_annotation(event):
    # 当鼠标在坐标轴内时更新注释内容
    if event.inaxes == plt.gca():
        # 获取当前鼠标位置的原始特征值
        age = event.xdata
        salary = event.ydata
        # 格式化显示文本(保留1位小数的年龄,整数的薪资)
        text_content = f"年龄: {age:.1f}\n薪资: {int(salary):,}"
        annot.set_text(text_content)
        annot.xy = (age, salary)
        annot.set_visible(True)
        plt.draw()
    else:
        # 鼠标移出坐标轴时隐藏注释
        if annot.get_visible():
            annot.set_visible(False)
            plt.draw()

# 绑定鼠标移动事件到更新函数
plt.gcf().canvas.mpl_connect("motion_notify_event", update_annotation)

plt.show()

现在运行代码,当你把鼠标移到图表上时,左下角就会显示当前位置对应的年龄和薪资数值,完全符合你的要求。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.20 09:20:42