JAX新手求教:如何修改线性回归代码以读取波士顿房价数据集
适配波士顿房价数据集修改JAX线性回归代码的方法
嘿,我来帮你搞定这个问题!你的代码其实已经导入了波士顿房价数据集,但后面不小心用手动定义的数组把它覆盖了,只需要做几个小调整就能切换到真实数据集上运行,具体修改点如下:
- 删掉手动定义的X和y:把代码里那段手动创建特征矩阵和目标向量的代码完全移除,保留前面从
load_boston()加载的X和y。 - 调整权重w的初始化维度:波士顿数据集每个样本有13个特征,原来的
w = np.zeros((2, 1))要改成w = np.zeros((13, 1)),这样矩阵乘法X.dot(w)的维度才能匹配(X是(n_samples,13),w是(13,1),相乘后得到符合要求的(n_samples,1)预测值)。 - 统一目标向量的形状(推荐):加载的波士顿房价目标值
y是一维数组(形状为(n_samples,)),而你的损失函数里预测值y_hat是二维的,虽然JAX会自动处理维度广播,但为了避免潜在的混淆,建议把y转换成二维:y = np.array(boston.target).reshape(-1, 1)。 - (可选)清理冗余导入:你代码里导入了
sklearn.linear_model as sk但没用到,要是不需要的话可以删掉,让代码更清爽。
修改后的完整代码示例
import jax.numpy as np from jax import grad, jit from sklearn.datasets import load_boston # 加载波士顿房价数据集 boston = load_boston() X = np.array(boston.data) # 将目标向量转为二维,匹配预测值维度 y = np.array(boston.target).reshape(-1, 1) def J(X, w, b, y): """线性回归的损失函数,模型的前向传播过程。 参数: X: 特征矩阵 w: 权重(列向量) b: 偏置 y: 目标向量 返回: 标量: 当前解的损失值 """ y_hat = X.dot(w) + b # 计算预测值 return ((y_hat - y)**2).mean() # 返回损失值 learning_rate = 0.01 # 初始化权重:13个特征对应13个权重参数 w = np.zeros((13, 1)) b = 0. # 验证梯度计算效率 %timeit grad(J, argnums=1)(X, w, b, y) %timeit grad(J, argnums=2)(X, w, b, y) # 迭代训练 for i in range(100): w -= learning_rate * grad(J, argnums=1)(X, w, b, y) b -= learning_rate * grad(J, argnums=2)(X, w, b, y) if i % 10 == 0: print(f"迭代次数 {i}, 损失值: {J(X, w, b, y)}")
内容的提问来源于stack exchange,提问作者Simon Templar
相关产品推荐
相关产品推荐

