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

如何在Google JAX中对自定义卷积包装器正确使用grad梯度计算功能

问题原因
  • 传参逻辑错误:JAX的grad默认对函数的第一个输入参数求梯度,你调用grad(conv)(y, k)时,将输出特征维度为64的特征图y作为输入传入conv函数,而卷积核k的输入特征维度仅为3,二者维度不匹配触发了卷积算子的参数校验错误,就是你看到的报错信息。
  • 梯度计算逻辑不合法:grad要求目标函数的输出必须是标量才能计算梯度,你直接对输出为4维张量的conv函数求梯度,即使维度匹配也会报错。
解决方案

首先构造一个输出为标量的损失函数,再根据你需要求导的参数指定argnums参数调用grad,示例代码如下:

init = normal()
rng = PRNGKey(42)
x = init(rng, [128, 3, 224, 224])
k = init(rng, [64, 3, 3, 3])

# 定义标量损失,比如对卷积输出的所有元素求和
def loss(x, weight):
    return conv(x, weight).sum()

# 对卷积核weight(也就是函数的第二个输入参数,索引为1)求梯度
k_grad = grad(loss, argnums=1)(x, k)

# 如果需要同时对输入x和卷积核k求梯度,写法如下
x_grad, k_grad = grad(loss, argnums=(0, 1))(x, k)

可选优化(避免后续使用出错)

你当前的conv包装器存在两个小问题,可以提前修正:

  1. dilation参数赋值错误:普通卷积的空洞率是作用在卷积核上的,对应lax.conv_general_dilated的rhs_dilation参数,lhs_dilation仅用于转置卷积场景,赋值错误会导致带空洞的卷积计算结果异常。
  2. 不支持分组卷积:你写死了feature_group_count=1,没有用到函数传入的groups参数,无法实现分组卷积功能。

修正后的conv函数返回部分代码如下:

return lax.conv_general_dilated(
    lhs=input,
    rhs=weight,
    window_strides=stride,
    padding=padding,
    rhs_dilation=dilation,
    dimension_numbers=torch_dims[n],
    feature_group_count=groups,
    batch_group_count=1,
    precision=None,
    preferred_element_type=None
)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.30 09:36:03