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

MNIST高斯生成模型不同正则化参数c预测误差一致问题排查

代码问题排查
  • 未定义核心计算变量score:你在计算predictions时调用的score没有在函数内定义,用的是全局环境中固定的旧值,所以不管c怎么调整,预测结果完全不会变,这是不同参数误差一致的核心原因。score本质是每个样本属于10个类别的对数后验概率,需要你用拟合得到的mu、sigma、pi结合multivariate_normal.logpdf计算,还要加上先验pi的对数。
  • 函数功能设计不符合要求:任务要求fit_generative_model仅负责拟合给定训练集的高斯模型参数,返回pi、mu、sigma即可。你现在把c值遍历、预测、误差计算的逻辑全部耦合在fit函数里,逻辑完全混乱。
  • 遗漏先验概率pi的计算:你只声明了pi数组,没有进行赋值,pi是每个类别的先验概率,计算公式为pi[label] = np.sum(indices) / len(y),生成模型的后验概率计算必须包含先验项。
  • 训练测试流程完全颠倒:你直接把测试集传入fit函数拟合参数,相当于用测试集数据训练模型,完全违背了机器学习的基本流程。正确流程是:用训练集拟合得到模型参数,再用参数对测试集样本做预测,最后对比测试集真实标签计算误差。
  • 函数无对应返回值:你调用函数时试图接收mu、sigma、pi三个返回值,但现有函数没有return语句,默认返回None,你能跑出结果进一步说明用的是全局环境的残留变量。
修正后的代码示例
# 符合要求的拟合函数,输入训练集和正则化系数c,返回对应模型参数
def fit_generative_model(x, y, c):
    k = 10
    d = x.shape[1]
    mu = np.zeros((k, d))
    sigma = np.zeros((k, d, d))
    pi = np.zeros(k)
    total_num = len(y)
    for label in range(k):
        indices = (y == label)
        class_x = x[indices]
        mu[label] = np.mean(class_x, axis=0)
        sigma[label] = np.cov(class_x, rowvar=False, bias=True) + c * np.eye(d)
        pi[label] = len(class_x) / total_num
    return mu, sigma, pi

# 预测函数,输入模型参数和待预测样本,返回预测标签
def predict(mu, sigma, pi, x):
    k = len(pi)
    score = np.zeros((len(x), k))
    for label in range(k):
        # 计算每个类的对数似然 + 对数先验
        score[:, label] = multivariate_normal.logpdf(x, mean=mu[label], cov=sigma[label]) + np.log(pi[label])
    return np.argmax(score, axis=1)

# 测试不同c的效果
for c in [20, 200, 4000]:
    # 用训练集拟合参数
    mu, sigma, pi = fit_generative_model(train_data, train_labels, c)
    # 预测测试集
    test_pred = predict(mu, sigma, pi, test_data)
    # 计算测试误差
    test_error = np.sum(test_pred != test_labels) / len(test_labels)
    print(f"c={c},测试集错误率:{test_error * 100:.1f}%,错误数:{np.sum(test_pred != test_labels)}")

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.04 20:15:02