线性分类器散点图报错排查及超平面绘制方法求助
散点图ValueError解决
- 错误根源:
plt.scatter的第三个位置参数是s(控制点的大小),你传入的标签数组y包含-1,而点的大小不能为负数,因此触发ValueError。要按类别区分颜色,需用c参数指定颜色映射。 - 额外问题:代码中训练集变量名是小写
x,但训练和绘图时用了大写X,会导致NameError,需统一变量名。 - 修正后的绘图代码:
# 先统一训练集变量名为X np.random.seed(42) X = np.random.randn(100,2) y = np.concatenate([np.ones(50), -1*np.ones(50)]) # 绘制分类散点图 plt.scatter(X[:, 0], X[:, 1], c=y, cmap='bwr') plt.xlabel('Feature 1') plt.ylabel('Feature 2') plt.show()
绘制线性分类器的超平面
线性分类器的决策边界(超平面)满足公式:weights[0]*x₀ + weights[1]*x₁ + bias = 0,变形为x₁ = (-weights[0]*x₀ - bias)/weights[1]。通过生成x₀的取值范围,计算对应x₁即可绘制这条直线:
- 实现代码(放在训练完成后):
# 生成x₀的取值范围(覆盖数据集的特征范围) x0_range = np.linspace(X[:,0].min()-1, X[:,0].max()+1, 100) # 计算对应的x₁值 x1_values = (-weights[0]*x0_range - bias)/weights[1] # 绘制散点图+决策边界 plt.scatter(X[:,0], X[:,1], c=y, cmap='bwr') plt.plot(x0_range, x1_values, 'k--', label='Decision Boundary') plt.xlabel('Feature 1') plt.ylabel('Feature 2') plt.legend() plt.show()
完整修正代码
整合所有修正后的可运行代码:
import numpy as np import matplotlib.pyplot as plt # 生成训练集 np.random.seed(42) X = np.random.randn(100,2) y = np.concatenate([np.ones(50), -1*np.ones(50)]) # 生成测试集 X_test = np.random.randn(50,2) y_test = np.concatenate([np.ones(25), -1*np.ones(25)]) # 线性分类器训练函数 def linear_classifier(X, y, learning_rate=0.01, num_epochs=100): num_features = X.shape[1] weights = np.zeros(num_features) bias = 0 for epoch in range(num_epochs): for i in range(X.shape[0]): linear_output = np.dot(X[i], weights) + bias y_pred = np.sign(linear_output) error = y[i] - y_pred weights += learning_rate * error * X[i] bias += learning_rate * error return weights, bias # 训练模型 weights, bias = linear_classifier(X, y) # 测试模型并计算准确率 linear_output = np.dot(X_test, weights) + bias y_pred = np.sign(linear_output) accuracy = np.mean(y_pred == y_test) print("Accuracy:", accuracy) # 绘制结果 x0_range = np.linspace(X[:,0].min()-1, X[:,0].max()+1, 100) x1_values = (-weights[0]*x0_range - bias)/weights[1] plt.scatter(X[:,0], X[:,1], c=y, cmap='bwr') plt.plot(x0_range, x1_values, 'k--', label='Decision Boundary') plt.xlabel('Feature 1') plt.ylabel('Feature 2') plt.title('Linear Classifier Decision Boundary') plt.legend() plt.show()
内容的提问来源于stack exchange,提问作者sarika
相关产品推荐
相关产品推荐

