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

手动实现Adam算法求解线性回归时参数发散问题排查求助

Adam算法实现线性回归参数发散问题排查

核心问题:梯度计算错误

线性回归的MSE损失函数为:
$$L = \frac{1}{N}\sum_{i=1}^N (y_i - \hat{y}_i)^2, \quad \hat{y}_i = w x_i + b$$
对应的偏导数应为:

  • 对w的偏导:单样本下为 $-2(y_i - w x_i - b)x_i$
  • 对b的偏导:单样本下为 $-2(y_i - w x_i - b)$

你的grad函数公式完全偏离MSE梯度的正确推导,这是参数发散的根本原因。修正后的梯度函数:

def grad(x, y, w, b, par):
    y_hat = w * x + b
    if par == "w":
        return -2 * x * (y - y_hat)
    if par == "b":
        return -2 * (y - y_hat)

Adam实现的其他问题

  1. 偏差修正逻辑错误
    Adam的偏差修正需要结合迭代步数的指数计算:
    $$\hat{m}_t = \frac{m_t}{1 - \beta_1^{t+1}}, \quad \hat{v}_t = \frac{v_t}{1 - \beta_2^{t+1}}$$
    你直接固定除以1 - beta1,未考虑迭代次数的影响,初始阶段修正不足会导致步幅过大,加速参数发散。需要维护全局迭代步数用于修正计算。

  2. 参数更新顺序错误
    你更新完w后,用更新后的w计算b的梯度,导致梯度基于最新参数而非当前迭代步的参数。正确做法是先计算当前步w和b的梯度,再同时更新两个参数。

  3. 小批量实现逻辑偏差
    当前代码生成batch后逐个样本更新参数,属于随机梯度下降(SGD)而非小批量Adam。小批量模式应先计算整个batch的梯度均值,再执行一次参数更新。

修正后的核心代码示例

import numpy as np
import math

def grad(x, y, w, b, par):
    y_hat = w * x + b
    if par == "w":
        return -2 * x * (y - y_hat)
    if par == "b":
        return -2 * (y - y_hat)

# 初始化参数与超参数
w = 0.0
b = 0.0
mw, mb = 0.0, 0.0  # 一阶矩估计
vw, vb = 0.0, 0.0  # 二阶矩估计
beta1, beta2 = 0.9, 0.999
alpha = 0.001
epsilon = 1e-8
epochs = 1000
batch_size = 32
n = len(x)  # 假设x为输入数据数组

global_t = 0  # 全局迭代步数,用于偏差修正
for epoch in range(epochs):
    arr = np.random.randint(0, n, batch_size)
    if epoch % 100 == 0:
        print(f"Epoch {epoch}: w={w:.4f}, b={b:.4f}")
    
    # 计算小批量梯度均值
    grad_w_sum, grad_b_sum = 0.0, 0.0
    for i in arr:
        grad_w_sum += grad(x[i], y[i], w, b, 'w')
        grad_b_sum += grad(x[i], y[i], w, b, 'b')
    grad_w = grad_w_sum / batch_size
    grad_b = grad_b_sum / batch_size
    
    global_t += 1
    # 更新矩估计
    mw = beta1 * mw + (1 - beta1) * grad_w
    vw = beta2 * vw + (1 - beta2) * (grad_w ** 2)
    mb = beta1 * mb + (1 - beta1) * grad_b
    vb = beta2 * vb + (1 - beta2) * (grad_b ** 2)
    
    # 偏差修正
    mmw = mw / (1 - math.pow(beta1, global_t))
    vvw = vw / (1 - math.pow(beta2, global_t))
    mmb = mb / (1 - math.pow(beta1, global_t))
    vvb = vb / (1 - math.pow(beta2, global_t))
    
    # 更新参数
    w = w - (alpha * mmw) / (math.sqrt(vvw) + epsilon)
    b = b - (alpha * mmb) / (math.sqrt(vvb) + epsilon)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.29 10:05:28