感知机算法中误分类点的计算与可视化求助
感知机线性分类器:补全误分类点计算与可视化代码
问题描述
已完成感知机线性分类器的大部分Python代码编写,但在计算误分类点数量并将其可视化的环节遇到困难,需要补全代码实现该功能。原代码如下:
import numpy as np DATA = np.loadtxt("C:/Users/abish/OneDrive/Desktop/Machine Learning/data_Perceptron.txt") X = DATA[:, 0:2] Y = DATA[:, 2] import matplotlib.pyplot as plt fig, ax = plt.subplots(1, 1, figsize=(5, 5)) ax.scatter(X[:, 0], X[:, 1], c=Y) ax.set_title('ground truth', fontsize=20) plt.xlabel('X1') plt.ylabel('X2') plt.show() # Add a bias to the X vector X_bias = np.ones([X.shape[0], 3]) X_bias[:, 1:3] = X # Initialize weight vector with zeros w = np.zeros([3, 1]) # Define the activation function that returns either 1 or 0 def activation(x): return 1 if x >= 1 else 0 # A function to calculate the unit vector of our weights vector def calc_unit_vector(x): return x.transpose() / np.sqrt(x.transpose().dot(x)) # A function that returns values that lay on the hyperplane def calc_hyperplane(X, w): return np.ravel([-(w[0] + x * w[1]) / w[2] for x in X]) for _ in range(10): for i in range(X_bias.shape[0]): y = activation(w.transpose().dot(X_bias[i, :])) # Update weights w = w + ((Y[i] - y) * X_bias[i, :]).reshape(w.shape[0], 1) print('w0 = ', w[0]) print('w1 = ', w[1]) print('w2 = ', w[2]) # Calculate the class of the data points with the weight vector result = [w.transpose().dot(x) for x in X_bias] result_class = [activation(w.transpose().dot(x)) for x in X_bias] # Calculate unit vector w = calc_unit_vector(w).transpose() fig, ax = plt.subplots(1, 2, figsize=(15, 5)) ax[0].scatter(X[:, 0], X[:, 1], c=Y) ax[0].set_title('ground truth', fontsize=20) ax[1].scatter(X[:, 0], X[:, 1], c=result_class) ax[1].plot([-20, 20], calc_hyperplane([-20, 20], w), lw=3, c='red') ax[1].set_xlim(ax[0].get_xlim()) ax[1].set_ylim(ax[0].get_ylim()) ax[1].set_yticks([]) ax[1].set_title('Perceptron classification with hyperplane', fontsize=20) plt.show() # calculate misclassified points misclassified = fig, ax = plt.subplots(1, 1, figsize=(5, 5)) ax.scatter(X[:, 0], X[:, 1], c=misclassified) ax.set_title('misclassified points', fontsize=20) plt.xlabel('X1') plt.ylabel('X2') plt.show()
解决方案:补全误分类点计算与可视化代码
以下是补全后的完整代码,优化了误分类点的识别逻辑和可视化效果:
import numpy as np DATA = np.loadtxt("C:/Users/abish/OneDrive/Desktop/Machine Learning/data_Perceptron.txt") X = DATA[:, 0:2] Y = DATA[:, 2] import matplotlib.pyplot as plt fig, ax = plt.subplots(1, 1, figsize=(5, 5)) ax.scatter(X[:, 0], X[:, 1], c=Y) ax.set_title('ground truth', fontsize=20) plt.xlabel('X1') plt.ylabel('X2') plt.show() # Add a bias to the X vector X_bias = np.ones([X.shape[0], 3]) X_bias[:, 1:3] = X # Initialize weight vector with zeros w = np.zeros([3, 1]) # Define the activation function that returns either 1 or 0 def activation(x): return 1 if x >= 1 else 0 # A function to calculate the unit vector of our weights vector def calc_unit_vector(x): return x.transpose() / np.sqrt(x.transpose().dot(x)) # A function that returns values that lay on the hyperplane def calc_hyperplane(X, w): return np.ravel([-(w[0] + x * w[1]) / w[2] for x in X]) for _ in range(10): for i in range(X_bias.shape[0]): y = activation(w.transpose().dot(X_bias[i, :])) # Update weights w = w + ((Y[i] - y) * X_bias[i, :]).reshape(w.shape[0], 1) print('w0 = ', w[0]) print('w1 = ', w[1]) print('w2 = ', w[2]) # Calculate the class of the data points with the weight vector result = [w.transpose().dot(x) for x in X_bias] result_class = [activation(w.transpose().dot(x)) for x in X_bias] # Calculate unit vector w = calc_unit_vector(w).transpose() fig, ax = plt.subplots(1, 2, figsize=(15, 5)) ax[0].scatter(X[:, 0], X[:, 1], c=Y) ax[0].set_title('ground truth', fontsize=20) ax[1].scatter(X[:, 0], X[:, 1], c=result_class) ax[1].plot([-20, 20], calc_hyperplane([-20, 20], w), lw=3, c='red') ax[1].set_xlim(ax[0].get_xlim()) ax[1].set_ylim(ax[0].get_ylim()) ax[1].set_yticks([]) ax[1].set_title('Perceptron classification with hyperplane', fontsize=20) plt.show() # ------------------- 补全部分 ------------------- # 将预测结果转为numpy数组,方便与真实标签对比 result_class_np = np.array(result_class) # 生成误分类标记:1表示误分类,0表示分类正确 misclassified = (result_class_np != Y).astype(int) # 统计误分类点数量 misclassified_count = np.sum(misclassified) print(f"误分类点数量:{misclassified_count}") # 可视化:用红色高亮误分类点,灰色展示正确分类点,叠加分类超平面 fig, ax = plt.subplots(1, 1, figsize=(5, 5)) # 绘制正确分类的点 ax.scatter(X[misclassified == 0, 0], X[misclassified == 0, 1], c='gray', label='正确分类') # 绘制误分类的点(放大尺寸增强辨识度) ax.scatter(X[misclassified == 1, 0], X[misclassified == 1, 1], c='red', s=80, label='误分类') # 绘制分类超平面 ax.plot([-20, 20], calc_hyperplane([-20, 20], w), lw=3, c='blue', label='分类超平面') ax.set_title('misclassified points', fontsize=20) plt.xlabel('X1') plt.ylabel('X2') ax.legend() plt.show()
关键逻辑说明
- 误分类点识别:通过
result_class_np != Y生成布尔数组,直接标记每个样本的分类正误,再转为整数类型便于后续处理。 - 数量统计:利用
np.sum()快速统计误分类点的总数。 - 可视化优化:通过分层绘制不同类别的点,并用颜色和尺寸区分正误分类,叠加超平面让结果更直观。
内容的提问来源于stack exchange,提问作者Abishek Phuyal
相关产品推荐
相关产品推荐

