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

感知机(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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.10 11:05:14