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

JAX:能否通过字典而非参数索引指定自动微分变量?

用JAX实现基于字典参数的梯度下降

当然可行,JAX本身不直接支持按参数名称指定求导目标,但可以通过包装函数或重构损失函数的方式实现字典传参+按键求导的需求,下面是两种更简洁的实现方案:

方案一:重构损失函数为字典参数输入

直接把需要优化的参数(w、b)打包成字典,让损失函数接收字典形式的参数,这样可以直接对整个参数字典求导,再按键提取梯度更新:

import numpy as np
import jax.numpy as jnp
from jax import grad

# 数据准备
X = np.array([
                 [4., 7.],
                 [1., 8.],
                 [-5., -6.],
                 [3., -1.],
                 [0., 9.]
             ])
y = np.array([
                 [37.],
                 [24.],
                 [-34.], 
                 [16.],
                 [21.]
             ])
learning_rate = 0.01

# 初始化参数字典
params = {
    'w': jnp.zeros((2, 1)),
    'b': jnp.array(0.)
}

# 重构损失函数,接收参数字典、X、y
def J_dict(params, X, y):
    y_hat = X.dot(params['w']) + params['b']
    return ((y_hat - y)**2).mean()

# 梯度下降循环
for i in range(100):
    # 计算整个参数字典的梯度
    grads = grad(J_dict)(params, X, y)
    # 按键更新参数
    params['w'] = params['w'] - learning_rate * grads['w']
    params['b'] = params['b'] - learning_rate * grads['b']

这种方式最贴合你的需求,无需手动管理参数位置,直接通过字典键操作即可。

方案二:包装原函数适配字典传参

如果不想修改原有损失函数J(X, w, b, y),可以写一个包装函数,把字典参数转换成位置参数传给原函数,再针对目标键对应的位置求导:

import numpy as np
import jax.numpy as jnp
from jax import grad

# 原损失函数保持不变
def J(X, w, b, y):
    y_hat = X.dot(w) + b
    return ((y_hat - y)**2).mean()

# 数据与参数初始化
X = np.array([[4.,7.],[1.,8.],[-5.,-6.],[3.,-1.],[0.,9.]])
y = np.array([[37.],[24.],[-34.],[16.],[21.]])
learning_rate = 0.01

arg_dict = {
    'X': jnp.array(X),
    'w': jnp.zeros((2, 1)),
    'b': jnp.array(0.),
    'y': jnp.array(y)
}

# 参数名到位置的映射(根据原函数的参数顺序)
param_order = ['X', 'w', 'b', 'y']
name_to_idx = {name: idx for idx, name in enumerate(param_order)}

# 包装函数:接收字典,转成位置参数调用J
def J_wrapper(args_dict):
    return J(*[args_dict[name] for name in param_order])

# 梯度下降循环
for i in range(100):
    # 对w求导:指定argnums为w对应的位置
    grad_w = grad(J_wrapper, argnums=0)(arg_dict)['w']
    arg_dict['w'] = arg_dict['w'] - learning_rate * grad_w
    
    # 对b求导:指定argnums为b对应的位置
    grad_b = grad(J_wrapper, argnums=0)(arg_dict)['b']
    arg_dict['b'] = arg_dict['b'] - learning_rate * grad_b

优化你的临时解决方案

你的现有方案存在一个问题:var -= ...不会修改arg_dict里的参数值(因为numpy数组的赋值是副本操作),需要改为直接更新字典元素,同时建议使用JAX数组而非numpy数组以获得更好的兼容性:

import numpy as np
import jax.numpy as jnp
from jax import grad

def J(X, w, b, y):
    y_hat = X.dot(w) + b
    return ((y_hat - y)**2).mean()

X = np.array([[4.,7.],[1.,8.],[-5.,-6.],[3.,-1.],[0.,9.]])
y = np.array([[37.],[24.],[-34.],[16.],[21.]])
learning_rate = 0.01

arg_dict = {
    'X': jnp.array(X),
    'w': jnp.zeros((2, 1)),
    'b': jnp.array(0.),
    'y': jnp.array(y)
}

# 建立参数名到位置的映射(必须和J的参数顺序一致)
param_order = list(arg_dict.keys())
name_to_idx = {name: idx for idx, name in enumerate(param_order)}

for i in range(100):
    # 把字典转成位置参数列表
    arg_list = [arg_dict[name] for name in param_order]
    for name in ['w', 'b']:  # 只对需要优化的参数更新
        idx = name_to_idx[name]
        # 计算梯度并更新字典中的参数
        grad_val = grad(J, argnums=idx)(*arg_list)
        arg_dict[name] = arg_dict[name] - learning_rate * grad_val

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.08 06:05:22