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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.29 13:32:31