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
相关产品推荐
相关产品推荐

