感知机(PLA)算法实现迭代次数异常问题排查求助
问题原因与修复方案
核心问题:数据生成函数的垂直直线处理错误
你的PLA迭代次数远超预期,主要原因是gen_data函数在生成分割线时,没有处理**两点x坐标相同(垂直直线)**的情况:
- 当选中的两个点x坐标相等时,计算斜率
m会出现除以0,得到无穷大inf - 后续标签生成逻辑
y >= m*x + c会失效,导致所有样本被错误标记为同一类(比如全-1) - 这种极端线性可分场景会让PLA需要大量迭代才能收敛,直接拉高了平均迭代次数
修复步骤
1. 修正数据生成函数
修改gen_data,单独处理垂直/水平直线的情况,避免除以0的错误:
def gen_data(N=10): size = (N, 2) data = np.random.uniform(-1, 1, size) # 随机选取两个不重复的点 point1, point2 = data[np.random.choice(data.shape[0], 2, replace=False), :] # 处理垂直直线(x=常数) if point2[0] == point1[0]: labels = np.array([+1 if x > point1[0] else -1 for x, y in data]) # 处理水平直线(y=常数) elif point2[1] == point1[1]: labels = np.array([+1 if y > point1[1] else -1 for x, y in data]) # 普通斜线情况 else: m = (point2[1] - point1[1]) / (point2[0] - point1[0]) c = point1[1] - m * point1[0] # 用>替代>=,避免样本刚好落在直线上的极端情况(概率极低,但可避免) labels = np.array([+1 if y > m * x + c else -1 for x, y in data]) data = np.column_stack((data, labels)) return data, point1, point2
2. 优化PLA的sign函数(可选但更规范)
原sign函数返回0的情况,虽然不影响正确性,但会增加不必要的误分类判定。修改为标准PLA的符号函数逻辑:
def sign(self, z): return np.where(z > 0, 1, -1)
3. 修正迭代次数计数(可选)
原fit函数中,最后一次无错误的循环也会让count加1,导致计数多1次。可以调整计数时机:
def fit(self): while True: y_pred = self.predict(self.X) misclassified = np.where(y_pred != self.y)[0] if len(misclassified) == 0: break # 只有当需要更新时才计数 self.count += 1 idx = np.random.choice(misclassified) self.update_weight(idx)
验证效果
修改后,样本量为10时,PLA的平均迭代次数会回到15左右的正常范围。
内容的提问来源于stack exchange,提问作者Yiffany
相关产品推荐
相关产品推荐

