如何在JAX中为自定义Module类计算权重梯度?
如何在JAX中为自定义类的权重计算梯度?
你当前代码的问题在于,jax.grad(criterion)(out, target)是对**损失函数的第一个输入(模型输出)**求导,而非对模型的权重参数求导,所以只会得到对应输出的单一梯度值,而非你需要的线性层权重梯度。
JAX是函数式编程框架,和PyTorch的面向对象式自动微分逻辑不同:JAX的自动微分只能追踪函数输入参数的梯度,无法直接追踪类实例属性的梯度。要解决这个问题,需要把模型参数从类中抽离,作为显式的函数输入参与计算流程。
修改后的完整实现
import jax import jax.numpy as jnp from typing import Dict class Module: def __init__(self) -> None: pass def __call__(self, params: Dict, inputs: jax.Array): return self.forward(params, inputs) class Linear(Module): def __init__(self, key: jax.Array, in_features: int, out_features: int) -> None: super().__init__() self.in_features = in_features self.out_features = out_features def init_params(self, key: jax.Array) -> Dict: key_w, key_b = jax.random.split(key=key, num=2) return { "weights": jax.random.normal(key=key_w, shape=(self.out_features, self.in_features)), "biases": jax.random.normal(key=key_b, shape=(self.out_features,)) } def forward(self, params: Dict, inputs: jax.Array) -> jax.Array: out = jnp.dot(params["weights"], inputs) + params["biases"] return out class Activation(Module): def __init__(self) -> None: super().__init__() def init_params(self) -> Dict: # 激活层无参数,返回空字典 return {} def forward(self, params: Dict, inputs: jax.Array) -> jax.Array: return jax.nn.sigmoid(inputs) class Model(Module): def __init__(self, key: jax.Array, in_features: int, out_features: int) -> None: super().__init__() self.linear = Linear(key=key, in_features=in_features, out_features=out_features) self.activation = Activation() def init_params(self, key: jax.Array) -> Dict: key_linear, _ = jax.random.split(key) return { "linear": self.linear.init_params(key_linear), "activation": self.activation.init_params() } def forward(self, params: Dict, inputs: jax.Array) -> jax.Array: out = self.linear(params["linear"], inputs) out = self.activation(params["activation"], out) return out def criterion(params: Dict, model: Model, inputs: jax.Array, target: jax.Array): output = model(params, inputs) return ((target - output) ** 2).sum() if __name__ == "__main__": in_features: int = 4 out_features: int = 1 key = jax.random.PRNGKey(67) model = Model(key=key, in_features=in_features, out_features=out_features) # 初始化模型参数 params = model.init_params(key) key_data = jax.random.PRNGKey(68) data = jax.random.normal(key=key_data, shape=(in_features,)) target = jnp.array([2.0]) # 计算损失 loss = criterion(params, model, data, target) print(f"{loss = }") # 计算参数的梯度 grads = jax.grad(criterion)(params, model, data, target) # 打印线性层权重的梯度 print(f"线性层权重梯度:\n{grads['linear']['weights']}") print(f"线性层偏置梯度:\n{grads['linear']['biases']}")
关键改动说明
- 参数显式化:每个模块新增
init_params方法,返回该模块的参数字典;模型整体的参数是嵌套字典,用JAX的参数树(Parameter Tree)管理,JAX会自动处理嵌套结构的梯度计算。 - forward方法适配:所有
forward和__call__方法都接收params作为第一个参数,函数计算完全依赖输入参数而非类实例属性,符合JAX的函数式要求。 - 损失函数重构:新的
criterion函数接收参数、模型、输入和目标,内部调用模型计算输出并返回损失,确保jax.grad能追踪参数的梯度。 - 梯度计算:
jax.grad(criterion)返回的梯度结构与params完全一致,可直接取出线性层的权重和偏置梯度。
如果想更贴近PyTorch的使用体验,也可以使用flax(JAX生态的官方神经网络库),它封装了类似PyTorch的Module类,但底层仍遵循JAX的函数式逻辑。
内容的提问来源于stack exchange,提问作者Gilfoyle
相关产品推荐
相关产品推荐

