使用Python实现线性回归的代价函数和梯度下降算法
使用Python实现线性回归的代价函数和梯度下降算法
线性回归是一种常用的机器学习算法,用于预测连续值。它涉及找到一条最适合数据的直线,通过最小化代价函数来实现。
本文将引导你使用Python实现线性回归的代价函数和梯度下降算法,并附带代码解释和示例。
1. 代价函数
代价函数衡量线性回归模型预测值与实际值之间的差异。我们的目标是找到使代价函数最小化的模型参数。
以下是线性回归的代价函数公式:
J(θ) = (1 / 2m) * Σ(hθ(x^(i)) - y^(i))^2
其中:
- J(θ) 是代价函数
- m 是训练样本的数量
- hθ(x^(i)) 是由参数 θ 定义的模型对第 i 个样本的预测
- y^(i) 是第 i 个样本的实际值
以下是使用Python实现代价函数的代码:
import numpy as np
import pandas as pd
# 计算代价函数 J(θ)
def computeCost(X, y, theta):
# X、y、theta 与数据预处理参数保持一致
# 返回代价函数的值(参考编程要求中的代价函数函数)
# ********** Begin **********
m = len(y)
predictions = X.dot(theta)
square_err = (predictions - y) ** 2
J = (1.0 / (2 * m)) * np.sum(square_err)
return J
# ********** End **********
2. 梯度下降
梯度下降是一种迭代算法,用于找到使代价函数最小化的模型参数。它通过在与代价函数梯度的相反方向上重复更新参数来实现这一点。
以下是梯度下降算法的公式:
θj := θj - α * (1 / m) * Σ(hθ(x^(i)) - y^(i)) * x^(i)_j
其中:
- θj 是第 j 个参数
- α 是学习率
- m 是训练样本的数量
- hθ(x^(i)) 是由参数 θ 定义的模型对第 i 个样本的预测
- y^(i) 是第 i 个样本的实际值
- x^(i)_j 是第 i 个样本中第 j 个特征的值
以下是使用Python实现梯度下降算法的代码:
# 批量梯度下降
# 返回参数 θ 的值和每次迭代后代价函数的值,梯度下降公式参考编程要求梯度下降公式
def gradientDescent(X, y, theta, alpha, epoch):
# X、y、theta 与数据预处理参数保持一致
# alpha: 学习率(取值:alpha = 0.01)
# epoch: 迭代次数(取值:epoch = 1000)
cost = np.zeros(epoch) # 初始化一个 ndarray ,包含每次 epoch 的 cost
# ********** Begin **********
m = len(y)
for i in range(epoch):
predictions = X.dot(theta)
theta = theta - alpha * (1.0 / m) * X.T.dot(predictions - y)
cost[i] = computeCost(X, y, theta)
# ********** End **********
return theta, cost
3. 示例
以下是如何使用这些函数执行线性回归的示例:
# 导入所需库
import matplotlib.pyplot as plt
# 加载数据
data = pd.read_csv('housing_data.csv')
X = data.iloc[:, :-1].values
y = data.iloc[:, -1].values
# 添加一列 1 到 X 以表示截距项
m = len(y)
X = np.concatenate((np.ones((m, 1)), X), axis=1)
# 初始化参数
theta = np.zeros(X.shape[1])
# 设置超参数
alpha = 0.01
epoch = 1000
# 使用梯度下降训练模型
theta, cost = gradientDescent(X, y, theta, alpha, epoch)
# 打印最终参数
print('最终参数:', theta)
# 绘制代价函数
plt.plot(np.arange(epoch), cost)
plt.xlabel('迭代次数')
plt.ylabel('代价函数')
plt.show()
结论
本文介绍了如何使用Python实现线性回归的代价函数和梯度下降算法。这些概念是理解和实现更高级的机器学习算法的基础。
原文地址: https://www.cveoy.top/t/topic/jULM 著作权归作者所有。请勿转载和采集!