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

如何在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']}")

关键改动说明

  1. 参数显式化:每个模块新增init_params方法,返回该模块的参数字典;模型整体的参数是嵌套字典,用JAX的参数树(Parameter Tree)管理,JAX会自动处理嵌套结构的梯度计算。
  2. forward方法适配:所有forward和__call__方法都接收params作为第一个参数,函数计算完全依赖输入参数而非类实例属性,符合JAX的函数式要求。
  3. 损失函数重构:新的criterion函数接收参数、模型、输入和目标,内部调用模型计算输出并返回损失,确保jax.grad能追踪参数的梯度。
  4. 梯度计算:jax.grad(criterion)返回的梯度结构与params完全一致,可直接取出线性层的权重和偏置梯度。

如果想更贴近PyTorch的使用体验,也可以使用flax(JAX生态的官方神经网络库),它封装了类似PyTorch的Module类,但底层仍遵循JAX的函数式逻辑。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.04 12:04:51