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

基于Numpy的梯度下降函数报错问题求助

梯度下降函数修正:TypeError及逻辑缺陷解决

问题描述

我编写了一个gradientDescent函数,输入为train_X、train_y、test_X、test_y、L(学习率)、num_iter(迭代次数),需返回最优参数及训练集、测试集的RMSE损失历史。已实现RMSE计算函数,但梯度下降算法存在逻辑缺陷,运行时抛出TypeError: unsupported operand type(s) for -: 'NoneType' and 'float'错误,请求修正代码逻辑并解决该错误。

已实现的RMSE函数

import numpy as np
import math

def RMSEFunction(X, theta, y):
    loss = None
    predicted_y = np.dot(X, theta)
    diff = np.subtract(predicted_y, y)
    diff_sq = np.square(diff)
    MSE = np.mean(diff_sq)
    loss = math.sqrt(MSE)
    return loss

错误分析与修正点

  1. TypeError根源:opt_theta初始化为None,第一次迭代中执行opt_theta -= ...时,试图让None与数值做减法运算,直接触发类型错误。
  2. 参数更新对象错误:迭代中应更新模型参数theta,而非用于保存最优参数的opt_theta。
  3. 迭代逻辑顺序混乱:原代码先错误更新opt_theta,再记录未更新的theta对应的损失,不符合梯度下降的迭代流程。

修正后的梯度下降函数

def gradientDescent(train_X, train_y, test_X, test_y, L, num_iter):  
    N_train, D = train_X.shape  # 训练样本数、特征数
    theta = np.zeros((D, 1))    # 初始化模型参数
    opt_theta = np.copy(theta)  # 初始化最优参数为初始参数副本

    # 初始化损失历史记录:行0为训练集RMSE,行1为测试集RMSE
    loss_history = np.zeros((2, num_iter))
    # 初始化最优测试集损失为初始参数对应的测试集RMSE
    test_loss = RMSEFunction(test_X, theta, test_y)
  
    # 梯度下降迭代
    for i in range(num_iter): 
        # 计算当前参数下的训练集预测值
        predicted_y = np.dot(train_X, theta)
        # 计算预测误差
        error = np.subtract(predicted_y, train_y)
        # 计算MSE的梯度(除以样本数得到均值梯度)
        gradient = np.dot(train_X.T, error) / N_train
        # 更新模型参数
        theta -= L * gradient

        # 记录当前参数对应的训练集、测试集RMSE
        train_rmse = RMSEFunction(train_X, theta, train_y)
        test_rmse = RMSEFunction(test_X, theta, test_y)
        loss_history[0, i] = train_rmse
        loss_history[1, i] = test_rmse
  
        # 更新最优参数(基于测试集损失最小化)
        if test_rmse < test_loss:
            test_loss = test_rmse
            opt_theta = np.copy(theta)
  
    return [opt_theta, loss_history]

额外说明

  • 修正后的代码严格遵循梯度下降流程:计算预测值→误差→梯度→更新参数→记录损失→更新最优参数。
  • 梯度计算时直接除以N_train,让梯度为均值梯度,与学习率的配合更符合标准线性回归的梯度下降公式。
  • 确保opt_theta始终持有合法的参数数组,避免类型错误。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.18 14:31:51